# %% [code]
import enum
import os
import warnings
from pathlib import Path
from typing import Optional

import joblib
import numpy as np
import pandas as pd
import pydicom
from typing_extensions import Annotated, Self

try:
    import pandera as pa
    from pandera.typing import Category, DataFrame, Object, Series
except ModuleNotFoundError:
    os.system("pip install pandera -q")
    import pandera as pa
    from pandera.typing import Category, DataFrame, Object, Series


class StrEnum(str, enum.Enum):

    """
    StrEnum is a Python ``enum.Enum`` that inherits from ``str``. The default
    ``auto()`` behavior uses the lowered member name as its value.

    Example usage::


        class Example(StrEnum):

            UPPER_CASE = auto()
            lower_case = auto()
            MixedCase = auto()


        assert Example.UPPER_CASE == "UPPER_CASE"
        assert Example.lower_case == "lower_case"
        assert Example.MixedCase == "MixedCase"

    """

    def __new__(cls, value: str | enum.auto, *args, **kwargs):
        """Generate a new member of the enumeration.

        Args:
            value: The value of the member.

        Raises:
            TypeError: If the value of the member is not a `str` or `enum.auto`.
        """
        if not isinstance(value, (str, enum.auto)):
            msg = f"Values of StrEnums must be strings: {value!r} is a {type(value)}"
            raise TypeError(msg)
        return super().__new__(cls, value, *args, **kwargs)

    def __str__(self) -> str:
        """Get string representation of member.

        Returns:
            The member represented as a str.
        """
        return str(self.value)

    @staticmethod
    def _generate_next_value_(name: str, *_) -> str:
        """Defines next member value when using `enum.auto()`.

        Args:
            name (str): The name of the member.

        Returns:
            The name of the member lowered.
        """
        return name.lower()


class SeverityScore(enum.IntEnum):
    """Normal/Mild, Moderate, or Severe"""

    NORMAL_MILD = enum.auto()
    MODERATE = enum.auto()
    SEVERE = enum.auto()

    @classmethod
    def _label_mapping(cls) -> dict:
        return {
            "Normal/Mild": cls.NORMAL_MILD,
            "Moderate": cls.MODERATE,
            "Severe": cls.SEVERE,
        }

    @classmethod
    def _inverse_label_mapping(cls) -> dict:
        return {v: k for k, v in cls._label_mapping().items()}

    @property
    def label(self) -> str:
        return self._inverse_label_mapping()[self]

    @property
    def submission_label(self) -> str:
        return self.name.lower()

    @classmethod
    def from_label(cls, label: str) -> Self:
        return cls._label_mapping()[label]


SeverityScoreCategorical = Annotated[
    Category,
    [member.value for member in SeverityScore],
]


class TrainSchema(pa.DataFrameModel):
    study_id: Series[int] = pa.Field(unique=True)
    spinal_canal_stenosis_l1_l2: Series[SeverityScoreCategorical]
    spinal_canal_stenosis_l2_l3: Series[SeverityScoreCategorical]
    spinal_canal_stenosis_l3_l4: Series[SeverityScoreCategorical]
    spinal_canal_stenosis_l4_l5: Series[SeverityScoreCategorical]
    spinal_canal_stenosis_l5_s1: Series[SeverityScoreCategorical]
    left_neural_foraminal_narrowing_l1_l2: Series[SeverityScoreCategorical]
    left_neural_foraminal_narrowing_l2_l3: Series[SeverityScoreCategorical]
    left_neural_foraminal_narrowing_l3_l4: Series[SeverityScoreCategorical]
    left_neural_foraminal_narrowing_l4_l5: Series[SeverityScoreCategorical]
    left_neural_foraminal_narrowing_l5_s1: Series[SeverityScoreCategorical]
    right_neural_foraminal_narrowing_l1_l2: Series[SeverityScoreCategorical]
    right_neural_foraminal_narrowing_l2_l3: Series[SeverityScoreCategorical]
    right_neural_foraminal_narrowing_l3_l4: Series[SeverityScoreCategorical]
    right_neural_foraminal_narrowing_l4_l5: Series[SeverityScoreCategorical]
    right_neural_foraminal_narrowing_l5_s1: Series[SeverityScoreCategorical]
    left_subarticular_stenosis_l1_l2: Series[SeverityScoreCategorical]
    left_subarticular_stenosis_l2_l3: Series[SeverityScoreCategorical]
    left_subarticular_stenosis_l3_l4: Series[SeverityScoreCategorical]
    left_subarticular_stenosis_l4_l5: Series[SeverityScoreCategorical]
    left_subarticular_stenosis_l5_s1: Series[SeverityScoreCategorical]
    right_subarticular_stenosis_l1_l2: Series[SeverityScoreCategorical]
    right_subarticular_stenosis_l2_l3: Series[SeverityScoreCategorical]
    right_subarticular_stenosis_l3_l4: Series[SeverityScoreCategorical]
    right_subarticular_stenosis_l4_l5: Series[SeverityScoreCategorical]
    right_subarticular_stenosis_l5_s1: Series[SeverityScoreCategorical]


index_col = TrainSchema.study_id
targets = [
    TrainSchema.spinal_canal_stenosis_l1_l2,
    TrainSchema.spinal_canal_stenosis_l2_l3,
    TrainSchema.spinal_canal_stenosis_l3_l4,
    TrainSchema.spinal_canal_stenosis_l4_l5,
    TrainSchema.spinal_canal_stenosis_l5_s1,
    TrainSchema.left_neural_foraminal_narrowing_l1_l2,
    TrainSchema.left_neural_foraminal_narrowing_l2_l3,
    TrainSchema.left_neural_foraminal_narrowing_l3_l4,
    TrainSchema.left_neural_foraminal_narrowing_l4_l5,
    TrainSchema.left_neural_foraminal_narrowing_l5_s1,
    TrainSchema.right_neural_foraminal_narrowing_l1_l2,
    TrainSchema.right_neural_foraminal_narrowing_l2_l3,
    TrainSchema.right_neural_foraminal_narrowing_l3_l4,
    TrainSchema.right_neural_foraminal_narrowing_l4_l5,
    TrainSchema.right_neural_foraminal_narrowing_l5_s1,
    TrainSchema.left_subarticular_stenosis_l1_l2,
    TrainSchema.left_subarticular_stenosis_l2_l3,
    TrainSchema.left_subarticular_stenosis_l3_l4,
    TrainSchema.left_subarticular_stenosis_l4_l5,
    TrainSchema.left_subarticular_stenosis_l5_s1,
    TrainSchema.right_subarticular_stenosis_l1_l2,
    TrainSchema.right_subarticular_stenosis_l2_l3,
    TrainSchema.right_subarticular_stenosis_l3_l4,
    TrainSchema.right_subarticular_stenosis_l4_l5,
    TrainSchema.right_subarticular_stenosis_l5_s1,
]


class SeriesDescriptionsSchema(pa.DataFrameModel):
    series_id: Series[int] = pa.Field(unique=True)
    study_id: Series[int]
    series_description: Series[str]


class TrainLabelCoordinatesSchema(pa.DataFrameModel):
    study_id: Series[int]
    series_id: Series[int]
    instance_number: Series[int]
    condition: Series[str]
    level: Series[str]
    x: Series[float]
    y: Series[float]


class SubmissionSchema(pa.DataFrameModel):
    row_id: Series[str] = pa.Field(unique=True)
    normal_mild: Series[float] = pa.Field(in_range={"min_value": 0.0, "max_value": 1.0})
    moderate: Series[float] = pa.Field(in_range={"min_value": 0.0, "max_value": 1.0})
    severe: Series[float] = pa.Field(in_range={"min_value": 0.0, "max_value": 1.0})


class DICOMTagsSchema(pa.DataFrameModel):
    study_id: Series[int] = pa.Field(alias="study_id")
    series_id: Series[int] = pa.Field(alias="series_id")
    instance_id: Series[int] = pa.Field(alias="instance_id")
    bits_allocated: Series[int] = pa.Field(alias="BitsAllocated")
    bits_stored: Series[int] = pa.Field(alias="BitsStored")
    columns: Series[int] = pa.Field(alias="Columns")
    content_date: Series[str] = pa.Field(alias="ContentDate")
    content_time: Series[str] = pa.Field(alias="ContentTime")
    frame_of_reference_uid: Series[str] = pa.Field(alias="FrameOfReferenceUID")
    high_bit: Series[int] = pa.Field(alias="HighBit")
    image_orientation_patient: Series[Object] = pa.Field(
        alias="ImageOrientationPatient",
    )
    image_position_patient: Series[Object] = pa.Field(alias="ImagePositionPatient")
    instance_number: Series[int] = pa.Field(alias="InstanceNumber")
    patient_id: Series[str] = pa.Field(alias="PatientID")
    patient_position: Series[Object] = pa.Field(alias="PatientPosition")
    photometric_interpretation: Series[str] = pa.Field(
        alias="PhotometricInterpretation",
    )
    pixel_representation: Series[int] = pa.Field(alias="PixelRepresentation")
    pixel_spacing: Series[Object] = pa.Field(alias="PixelSpacing")
    rows: Series[int] = pa.Field(alias="Rows")
    sop_instance_uid: Series[str] = pa.Field(alias="SOPInstanceUID")
    samples_per_pixel: Series[int] = pa.Field(alias="SamplesPerPixel")
    series_description: Series[str] = pa.Field(alias="SeriesDescription")
    series_instance_uid: Series[str] = pa.Field(alias="SeriesInstanceUID")
    slice_location: Series[float] = pa.Field(alias="SliceLocation")
    slice_thickness: Series[float] = pa.Field(alias="SliceThickness")
    spacing_between_slices: Series[float] = pa.Field(alias="SpacingBetweenSlices")
    study_instance_uid: Series[str] = pa.Field(alias="StudyInstanceUID")
    window_center: Series[float] = pa.Field(alias="WindowCenter")
    window_width: Series[float] = pa.Field(alias="WindowWidth")
    rescale_intercept: Series[float] = pa.Field(alias="RescaleIntercept", nullable=True)
    rescale_slope: Series[float] = pa.Field(alias="RescaleSlope", nullable=True)
    rescale_type: Optional[Series[str]] = pa.Field(alias="RescaleType", nullable=True)


class FileSystemNode:
    def __init__(self, parent: Optional["FileSystemNode"] = None):
        self.parent = parent

    @property
    def platform_is_kaggle(self):
        return os.environ["PWD"] == "/kaggle/working"


class CoreData(FileSystemNode):
    def __init__(self, parent: FileSystemNode):
        super().__init__(parent)

        if self.platform_is_kaggle:
            root_dir = (
                "/kaggle/input/rsna-2024-lumbar-spine-degenerative-classification"
            )
        else:
            root_dir = str(Path(__file__).parent.parent / "data")

        self.root_dir = Path(root_dir)

        self.train_imgs_dir = self.root_dir / "train_images"
        self.test_imgs_dir = self.root_dir / "test_images"

        self.sample_submission_path = self.root_dir / "sample_submission.csv"
        self.train_series_descriptions_path = (
            self.root_dir / "train_series_descriptions.csv"
        )
        self.train_label_coordinates_path = (
            self.root_dir / "train_label_coordinates.csv"
        )
        self.test_series_descriptions_path = (
            self.root_dir / "test_series_descriptions.csv"
        )
        self.train_path = self.root_dir / "train.csv"

    def get_study_path(self, study_id: int, *, train: bool) -> Path:
        dir_path = self.train_imgs_dir if train else self.test_imgs_dir
        return dir_path / str(study_id)

    def get_series_path(self, study_id: int, series_id: int, *, train: bool) -> Path:
        return self.get_study_path(study_id, train=train) / str(series_id)

    def get_instance_paths(
        self,
        study_id: int,
        series_id: int,
        *,
        train: bool,
    ) -> list[Path]:
        path = self.get_series_path(study_id, series_id, train=train)
        return sorted(path.glob("*.dcm"), key=lambda x: x.stem.zfill(4))

    def get_study_ids(self, *, train: bool) -> list[int]:
        dir_path = self.train_imgs_dir if train else self.test_imgs_dir
        return sorted(int(path.stem) for path in dir_path.glob("*"))

    def get_series_ids(self, study_id: int, *, train: bool) -> list[int]:
        dir_path = self.get_study_path(study_id, train=train)
        return sorted(int(path.stem) for path in dir_path.glob("*"))

    def load_train(self) -> DataFrame[TrainSchema]:
        df = pd.read_csv(self.train_path)

        row_has_missing_mask = df.isna().any(axis=1)
        if row_has_missing_mask.any():
            warnings.warn(
                (
                    f"Train dataset has ({sum(row_has_missing_mask)} / {len(df)}) "
                    "rows with missing value(s). Removing."
                ),
                stacklevel=1,
            )
        df = df[~row_has_missing_mask]

        for target in targets:
            df[target] = df[target].apply(SeverityScore.from_label).astype("category")

        return DataFrame[TrainSchema](df)

    def load_train_series_descriptions(self) -> DataFrame[SeriesDescriptionsSchema]:
        return DataFrame[SeriesDescriptionsSchema](
            pd.read_csv(self.train_series_descriptions_path),
        )

    def load_test_series_descriptions(self) -> DataFrame[SeriesDescriptionsSchema]:
        return DataFrame[SeriesDescriptionsSchema](
            pd.read_csv(self.test_series_descriptions_path),
        )

    def load_train_label_coordinates(self) -> DataFrame[TrainLabelCoordinatesSchema]:
        return DataFrame[TrainLabelCoordinatesSchema](
            pd.read_csv(self.train_label_coordinates_path),
        )

    def load_sample_submission(self) -> DataFrame[SubmissionSchema]:
        return DataFrame[SubmissionSchema](pd.read_csv(self.sample_submission_path))

    def load_series(
        self,
        study_id: int,
        series_id: int,
        dicom_tags_df: pd.DataFrame,
        *,
        train: bool,
    ) -> tuple[np.ndarray, tuple[float, float, float]]:
        instance_paths = self.get_instance_paths(
            study_id=study_id,
            series_id=series_id,
            train=train,
        )
        series_dicom_tags_df = dicom_tags_df[
            (dicom_tags_df[DICOMTagsSchema.study_id] == study_id)
            & (dicom_tags_df[DICOMTagsSchema.series_id] == series_id)
        ]
        iop, (dx, dy), dz, pr, ba, bs = series_dicom_tags_df[
            [
                DICOMTagsSchema.image_orientation_patient,
                DICOMTagsSchema.pixel_spacing,
                DICOMTagsSchema.slice_thickness,
                DICOMTagsSchema.pixel_representation,
                DICOMTagsSchema.bits_allocated,
                DICOMTagsSchema.bits_stored,
            ]
        ].iloc[0]

        bit_shift = None
        if pr == 1:
            bit_shift = ba - bs

        @joblib.delayed
        def load_dcm(path, bit_shift):
            dcm = pydicom.dcmread(path)
            pixel_array = dcm.pixel_array
            if bit_shift is not None:
                dtype = pixel_array.dtype
                pixel_array = (pixel_array << bit_shift).astype(dtype) >> bit_shift
            return pixel_array.astype(np.float32)

        imgs = joblib.Parallel(n_jobs=-1)(
            load_dcm(path, bit_shift=bit_shift) for path in instance_paths
        )
        vol = np.dstack(imgs)
        imaging_axis = np.cross(iop[:3], iop[3:])
        distance_projection = np.dot(
            np.vstack(
                series_dicom_tags_df[DICOMTagsSchema.image_position_patient].to_numpy(),
            ),
            imaging_axis,
        )
        vol = vol[:, :, np.argsort(distance_projection)]
        vol = vol.transpose((1, 0, 2))
        spacing = (float(dx), float(dy), float(dz))
        return vol, spacing


class DICOMTagsData(FileSystemNode):
    def __init__(self, parent: FileSystemNode):
        super().__init__(parent)

        if self.platform_is_kaggle:
            root_dir = "/kaggle/input/rsna-2024-lumbar-spine-metadata"
        else:
            root_dir = str(
                Path(__file__).parent.parent
                / "data"
                / "rsna-2024-lumbar-spine-metadata",
            )

        self.root_dir = Path(root_dir)
        self.train_dicom_tags_path = self.root_dir / "train_metadata.parquet"
        self.test_dicom_tags_path = self.root_dir / "test_metadata.parquet"

    def load_dicom_tags(self, *, train: bool) -> DataFrame[DICOMTagsSchema]:
        path = self.train_dicom_tags_path if train else self.test_dicom_tags_path
        df = pd.read_parquet(path, engine="fastparquet")
        return DataFrame[DICOMTagsSchema](df)


class FileSystem(FileSystemNode):
    def __init__(self):
        super().__init__()
        self.core = CoreData(self)
        self.dicom_tags = DICOMTagsData(self)
