"""
Loader for Calcium & TTX/NBQX Datasets
=====================================================

This module allows loading multi-frame TIFF recordings from:
1. Calcium dyes: calbr, cal520, fluo4, fluo8, ef630
2. TTX/NBQX datasets: ttx, nbqxap5

Examples:
---------------
# Calcium dataset
stack = load_recording("C:/Dyes/calcium_dataset", "calbr", slice_num=1, recording=3)
print(stack.shape)  # (frames, height, width)

# TTX/NBQX dataset
stack = load_recording("C:/Dyes/calcium_dataset", "ttx", slice_num=1, recording=2)
print(stack.shape)
"""

import os
import sys
import subprocess

try:
    import tifffile
except ImportError:
    print("Installing required package: tifffile")
    subprocess.check_call([sys.executable, "-m", "pip", "install", "tifffile"])
    import tifffile


def load_recording(base_path, dataset, slice_num, recording):
    """
    Load a multi-frame TIFF recording for calcium dyes or TTX/NBQX datasets.

    Parameters you should select and replace for your case
    ----------
    base_path : str
        Path to the dataset folder.
    dataset : str
        One of ['calbr','cal520','fluo4','fluo8','ef630','ttx','nbqxap5']
    slice_num : int
        Slice number (1–7 for most, 1–4 for ttx, 1–7 for nbqxap5).
    recording : int
        Recording number (1–12, 1–9, or 1–4 depending on dataset).

    Returns
    -------
    numpy.ndarray : 3D array (frames, height, width)
    """

    
    dataset = dataset.lower().strip()

    calcium_dyes = ['calbr', 'cal520', 'fluo4', 'fluo8', 'ef630']
    ttx_nbqx = ['ttx', 'nbqxap5']
    all_datasets = calcium_dyes + ttx_nbqx

    if dataset not in all_datasets:
        raise ValueError(f"dataset must be one of {all_datasets}")

    # Slice limits
    if dataset == "ttx":
        if not (1 <= slice_num <= 4):
            raise ValueError("ttx slices must be 1–4")

    elif dataset == "nbqxap5":
        if not (1 <= slice_num <= 7):
            raise ValueError("nbqxap5 slices must be 1–7")

    else:  # Calcium dyes
        if not (1 <= slice_num <= 7):
            raise ValueError(f"{dataset} slices must be 1–7")

    # Recording limits
    if dataset in ['calbr', 'cal520', 'ef630']:
        max_rec = 12
    elif dataset in ['fluo4', 'fluo8']:
        max_rec = 9
    else:
        max_rec = 4  # ttx & nbqxap5 slice limits

    if not (1 <= recording <= max_rec):
        raise ValueError(f"{dataset} has only {max_rec} recordings per slice")

    
    file_name = f"{dataset}_s{slice_num:02d}_r{recording:02d}.tif"
    file_path = os.path.join(base_path, dataset, file_name)

    if not os.path.exists(file_path):
        raise FileNotFoundError(f"File not found:\n{file_path}")

    print(f"Loading file: {file_path}")
    return tifffile.imread(file_path)


# Test examples here you can run directly 
if __name__ == "__main__":
    base = r"C:/Dyes/calcium_dataset"

    # Calcium example
    stack1 = load_recording(base, "calbr", slice_num=1, recording=1)
    print("Calcium stack shape:", stack1.shape)

    # TTX example
    stack2 = load_recording(base, "ttx", slice_num=2, recording=2)
    print("TTX stack shape:", stack2.shape)
