#!/usr/bin/env python3
"""RSNA Knee 2026 - High-Fidelity Multi-Plane 2.5D Preprocessing & Caching Kernel (v3).

Processes raw DICOM MRI series into physically sorted, 140mm FOV cropped 2.5D triplets @ 336x336 px.
Saves compressed uint8 npz shards (~600 MB / chunk, ~10.5 GB total) to strictly respect Kaggle's 20GB disk limit.
Optimized for multi-core CPU execution on Kaggle instances.
"""

import os
import sys
import time
import gc
import json
from pathlib import Path
from concurrent.futures import ThreadPoolExecutor
from typing import Dict, List, Optional, Tuple

import numpy as np
import pandas as pd
import cv2
import pydicom

# Parameters
TARGET_SIZE = 336  # Divisible by 16 (DINOv3, 21 patches) and 14 (DINOv2, 24 patches)
TARGET_FOV_MM = 140.0
TRIPLETS_PER_SLOT = 3
N_WORKERS = min(os.cpu_count() or 4, 4)

CANONICAL_SLOTS = [
    "Sagittal_Fluid",
    "Sagittal_NonFluid",
    "Coronal_Fluid",
    "Coronal_NonFluid",
    "Axial_Fluid",
    "Axial_NonFluid",
]

def resolve_paths() -> Tuple[Path, Path, Path, Path]:
    """Dynamically resolve input paths on Kaggle and local environments."""
    print("Resolving environment paths...", flush=True)
    kaggle_input = Path("/kaggle/input")
    if kaggle_input.exists():
        print("Listing /kaggle/input contents:", flush=True)
        for item in kaggle_input.iterdir():
            print(f"  - {item.name}", flush=True)

    candidates = [
        Path("/kaggle/input/rsna-knee-abnormality-detection"),
        Path("/kaggle/input/competitions/rsna-knee-abnormality-detection"),
        Path("/kaggle/input"),
        Path("data"),
        Path("."),
    ]
    
    comp_root = None
    for c in candidates:
        if (c / "train_series.csv").exists():
            comp_root = c
            break
            
    if comp_root is None and kaggle_input.exists():
        found = list(kaggle_input.rglob("train_series.csv"))
        if found:
            comp_root = found[0].parent

    if comp_root is None:
        comp_root = Path("data")

    train_series_csv = comp_root / "train_series.csv"
    train_dcm_root = comp_root / "train_series"
    if not train_dcm_root.exists() and (comp_root / "train").exists():
        train_dcm_root = comp_root / "train"

    work_dir = Path("/kaggle/working") if Path("/kaggle/working").exists() else Path("data/engineered_volumes")
    work_dir.mkdir(parents=True, exist_ok=True)

    print(f"  comp_root: {comp_root}", flush=True)
    print(f"  train_series_csv: {train_series_csv} (exists: {train_series_csv.exists()})", flush=True)
    print(f"  train_dcm_root: {train_dcm_root} (exists: {train_dcm_root.exists()})", flush=True)
    print(f"  work_dir: {work_dir}", flush=True)

    return comp_root, train_series_csv, train_dcm_root, work_dir

def compute_slice_normal_key(ds) -> float:
    """Project 3D ImagePositionPatient onto the plane normal vector from ImageOrientationPatient."""
    try:
        if hasattr(ds, "ImageOrientationPatient") and hasattr(ds, "ImagePositionPatient"):
            orient = np.asarray(ds.ImageOrientationPatient, dtype=np.float64)
            normal = np.cross(orient[:3], orient[3:])
            pos = np.asarray(ds.ImagePositionPatient, dtype=np.float64)
            return float(np.dot(pos, normal))
        if hasattr(ds, "SliceLocation"):
            return float(ds.SliceLocation)
        if hasattr(ds, "InstanceNumber"):
            return float(ds.InstanceNumber)
    except Exception:
        pass
    return 0.0

def sort_series_slices(dcm_paths: List[Path]) -> List[Path]:
    """Sort DICOM paths along true physical slice normal vector."""
    records = []
    for i, path in enumerate(dcm_paths):
        try:
            ds = pydicom.dcmread(str(path), stop_before_pixels=True, force=True)
            key = compute_slice_normal_key(ds)
            records.append((key, path))
        except Exception:
            records.append((float(i), path))
    records.sort(key=lambda r: r[0])
    return [p for _, p in records]

def decode_dicom_slice(path: Path) -> Tuple[Optional[np.ndarray], Tuple[float, float], str]:
    """Decode raw DICOM pixel array, apply slope/intercept, and extract pixel spacing and laterality."""
    try:
        ds = pydicom.dcmread(str(path), force=True)
        arr = ds.pixel_array.astype(np.float32)
        slope = float(getattr(ds, "RescaleSlope", 1.0) or 1.0)
        intercept = float(getattr(ds, "RescaleIntercept", 0.0) or 0.0)
        arr = arr * slope + intercept
        if str(getattr(ds, "PhotometricInterpretation", "")).upper() == "MONOCHROME1":
            arr = float(np.nanmax(arr)) - arr
        spacing = tuple(float(x) for x in getattr(ds, "PixelSpacing", [0.5, 0.5]))
        lat = str(getattr(ds, "ImageLaterality", getattr(ds, "Laterality", ""))).upper()
        if np.isfinite(arr).any():
            return arr, spacing, lat
    except Exception:
        pass
    return None, (0.5, 0.5), ""

def crop_and_resize_physical_fov(img: np.ndarray, pixel_spacing: Tuple[float, float], target_fov_mm: float = TARGET_FOV_MM, size: int = TARGET_SIZE) -> np.ndarray:
    """Crop 140mm physical field centered on joint and resize to size x size."""
    h, w = img.shape
    sx, sy = pixel_spacing[0], pixel_spacing[1]
    crop_h = int(round(target_fov_mm / max(sy, 1e-4)))
    crop_w = int(round(target_fov_mm / max(sx, 1e-4)))
    if crop_h < h:
        sy_idx = (h - crop_h) // 2
        img = img[sy_idx : sy_idx + crop_h, :]
    if crop_w < w:
        sx_idx = (w - crop_w) // 2
        img = img[:, sx_idx : sx_idx + crop_w]
    return cv2.resize(img, (size, size), interpolation=cv2.INTER_AREA if img.shape[0] > size else cv2.INTER_LINEAR)

def process_slot_series(series_dir: Path, plane: str, triplet_count: int = TRIPLETS_PER_SLOT, size: int = TARGET_SIZE) -> Tuple[np.ndarray, bool]:
    """Extract physically aligned 2.5D triplets for one series in uint8 [0, 255]."""
    dcm_paths = list(series_dir.glob("*.dcm"))
    if not dcm_paths:
        return np.zeros((triplet_count, 3, size, size), dtype=np.uint8), False
        
    ordered_paths = sort_series_slices(dcm_paths)
    n_slices = len(ordered_paths)
    if n_slices == 0:
        return np.zeros((triplet_count, 3, size, size), dtype=np.uint8), False

    center_indices = np.linspace(0, n_slices - 1, triplet_count + 2)[1:-1].round().astype(int)
    needed_indices = sorted(list(set([max(0, c - 1) for c in center_indices] + list(center_indices) + [min(n_slices - 1, c + 1) for c in center_indices])))

    decoded = {}
    spacings = {}
    is_left = False
    for idx in needed_indices:
        arr, sp, lat = decode_dicom_slice(ordered_paths[idx])
        if arr is not None:
            decoded[idx] = arr
            spacings[idx] = sp
            if "L" in lat:
                is_left = True

    if not decoded:
        return np.zeros((triplet_count, 3, size, size), dtype=np.uint8), False

    # Series-level windowing (1.0% to 99.5%)
    pooled = np.concatenate([arr.ravel() for arr in decoded.values() if np.isfinite(arr).any()])
    lo = float(np.percentile(pooled, 1.0))
    hi = float(np.percentile(pooled, 99.5))
    if hi <= lo:
        hi = lo + 1e-4

    processed = {}
    for idx, raw in decoded.items():
        # Scale to [0, 255] uint8
        scaled = np.clip((raw - lo) / (hi - lo), 0.0, 1.0) * 255.0
        cropped = crop_and_resize_physical_fov(scaled, pixel_spacing=spacings.get(idx, (0.5, 0.5)), size=size)
        processed[idx] = np.clip(np.round(cropped), 0, 255).astype(np.uint8)

    triplets = []
    for c in center_indices:
        s_prev = processed.get(max(0, c - 1), np.zeros((size, size), dtype=np.uint8))
        s_curr = processed.get(c, np.zeros((size, size), dtype=np.uint8))
        s_next = processed.get(min(n_slices - 1, c + 1), np.zeros((size, size), dtype=np.uint8))
        triplet = np.stack([s_prev, s_curr, s_next], axis=0)
        
        # Laterality correction
        if is_left:
            if "sagittal" in plane.lower():
                triplet = triplet[::-1]
            elif "coronal" in plane.lower() or "axial" in plane.lower():
                triplet = np.flip(triplet, axis=-1)
        triplets.append(triplet)

    res = np.stack(triplets, axis=0)  # (triplet_count, 3, size, size) uint8
    return res, True

def determine_slot_name(plane: str, fluid: int) -> Optional[str]:
    p = plane.strip().capitalize()
    if p in ("Sagittal", "Coronal", "Axial"):
        return f"{p}_{'Fluid' if int(fluid) == 1 else 'NonFluid'}"
    return None

def process_single_study(args) -> Tuple[str, np.ndarray, np.ndarray]:
    study_id, series_records, dcm_root = args
    slots_map = {}
    for plane, fluid, s_id in series_records:
        s_name = determine_slot_name(plane, fluid)
        if s_name and s_name not in slots_map:
            slots_map[s_name] = s_id

    n_slots = len(CANONICAL_SLOTS)
    vol = np.zeros((n_slots, TRIPLETS_PER_SLOT, 3, TARGET_SIZE, TARGET_SIZE), dtype=np.uint8)
    mask = np.zeros((n_slots, TRIPLETS_PER_SLOT), dtype=bool)

    for s_idx, s_name in enumerate(CANONICAL_SLOTS):
        s_id = slots_map.get(s_name)
        if s_id:
            s_dir = dcm_root / str(study_id) / str(s_id)
            if s_dir.exists():
                plane = s_name.split("_")[0]
                t_arr, ok = process_slot_series(s_dir, plane)
                if ok:
                    vol[s_idx] = t_arr
                    mask[s_idx, :] = True
    return str(study_id), vol, mask

def main():
    print("=== RSNA Knee 2026 - High-Fidelity 2.5D Volume Preprocessing (v3) ===", flush=True)
    print(f"Target FOV: {TARGET_FOV_MM} mm | Target Resolution: {TARGET_SIZE}x{TARGET_SIZE} px", flush=True)
    print(f"Canonical Slots: {len(CANONICAL_SLOTS)} ({', '.join(CANONICAL_SLOTS)})", flush=True)
    print(f"Storage Format: Compressed uint8 npz shards (~600MB/chunk, ~10.5GB total)", flush=True)
    print(f"Using {N_WORKERS} CPU worker threads", flush=True)
    
    comp_root, train_series_csv, train_dcm_root, work_dir = resolve_paths()
    
    if not train_series_csv.exists():
        print(f"ERROR: train_series.csv not found at {train_series_csv}. Exiting.", flush=True)
        sys.exit(1)

    df_series = pd.read_csv(train_series_csv)
    studies = df_series["StudyInstanceUID"].unique()
    print(f"Found {len(studies)} unique studies in {train_series_csv}", flush=True)

    tasks = []
    for study_id, group in df_series.groupby("StudyInstanceUID"):
        records = [(row["Anatomical_Plane"], row["Fluid_Sensitive"], row["SeriesInstanceUID"]) for _, row in group.iterrows()]
        tasks.append((study_id, records, train_dcm_root))

    batch_size = 250
    total_batches = (len(tasks) + batch_size - 1) // batch_size
    print(f"Processing {len(tasks)} studies across {total_batches} batches...", flush=True)

    manifest_records = []
    start_time = time.time()

    for b_idx in range(total_batches):
        b_t0 = time.time()
        batch_tasks = tasks[b_idx * batch_size : (b_idx + 1) * batch_size]
        print(f"\n[Batch {b_idx + 1}/{total_batches}] Processing {len(batch_tasks)} studies...", flush=True)
        
        batch_ids = []
        batch_vols = []
        batch_masks = []
        
        with ThreadPoolExecutor(max_workers=N_WORKERS) as pool:
            results = list(pool.map(process_single_study, batch_tasks))
            
        for s_id, vol, mask in results:
            batch_ids.append(s_id)
            batch_vols.append(vol)
            batch_masks.append(mask)
            manifest_records.append({
                "StudyInstanceUID": s_id,
                "batch": b_idx,
                "valid_slots": int(mask[:, 0].sum()),
            })

        chunk_vols = np.stack(batch_vols, axis=0)   # (B, 6, 3, 3, 336, 336) uint8
        chunk_masks = np.stack(batch_masks, axis=0) # (B, 6, 3) bool
        
        # Save as single compressed npz archive per chunk
        chunk_npz_path = work_dir / f"train_volumes_chunk_{b_idx:03d}.npz"
        np.savez_compressed(
            chunk_npz_path,
            volumes=chunk_vols,
            masks=chunk_masks,
            study_ids=np.array(batch_ids, dtype=object),
        )
            
        elapsed_b = time.time() - b_t0
        total_elapsed = time.time() - start_time
        mb_size = os.path.getsize(chunk_npz_path) / (1024**2)
        print(f"  Saved {chunk_npz_path.name}: {mb_size:.1f} MB (compressed) in {elapsed_b:.1f}s. Total elapsed: {total_elapsed/60:.1f}m", flush=True)
        
        del chunk_vols, chunk_masks, batch_vols, batch_masks
        gc.collect()

    df_manifest = pd.DataFrame(manifest_records)
    df_manifest.to_csv(work_dir / "train_volumes_manifest.csv", index=False)
    print(f"\n>> All {len(tasks)} studies completed in {(time.time() - start_time)/60:.1f}m!", flush=True)
    print(f"Manifest saved to {work_dir / 'train_volumes_manifest.csv'}", flush=True)

if __name__ == "__main__":
    main()
