import json
import os
import random
import traceback

import numpy as np
from torch.utils.data import IterableDataset
from tqdm import tqdm
from dataclasses import dataclass
from enum import Enum

from oracle_dataset import get_sample_oracle_file_segment


class JSONLMemmap:
    def __init__(self, path, verbose=False):
        self.path = path
        self._index = self.get_line_index_map(verbose=verbose)

    def get_line_index_map(self, verbose=False):
        line_positions = [0]
        with open(self.path, "rb") as f:
            for line in tqdm(f, disable=not verbose, desc="Making memmap index"):
                line_positions.append(len(line) + line_positions[-1])
        return np.array(line_positions[:-1])

    def get_line_from_index(self, index: int):
        with open(self.path, "rb") as f:
            f.seek(self._index[index])
            return json.loads(f.readline())

    def __getitem__(self, index: int):
        return self.get_line_from_index(index)

    def __len__(self):
        return len(self._index)

    def __iter__(self):
        for i in range(len(self)):
            yield self[i]


@dataclass
class AudioConfig:
    sample_rate: int = 48000
    n_channels: int = 2
    duration_s: float = 1.0
    is_vae: bool = False
    is_mert: bool = False
    is_musicfm: bool = False


@dataclass
class SampleData:
    data_wav: np.ndarray  # Raw audio in 48kHz (channel, length)
    filepath: str
    start_s: float
    data_mert: np.ndarray | None = None  # MERT embedding
    data_musicfm: np.ndarray | None = None  # MusicFM3 embedding


class AudioLoaderDataset(IterableDataset):
    """Gets SampleData objects from the oracle dataset, which contain raw wav data"""

    def __init__(
        self,
        audio_cfg: AudioConfig,
        metas_path: str,
        split="train",
    ):
        self.audio_cfg = audio_cfg
        self.metas_path = metas_path
        self.split = split

        # load metas as a memmap to save memory
        self.metas = JSONLMemmap(self.metas_path, verbose=False)

        self.random_cache = []

    def __iter__(self):
        return self

    def pad_audio(self, track):
        assert len(track.shape) == 2
        if track.shape[1] < int(48000 * self.audio_cfg.duration_s):
            pad_length = int(48000 * self.audio_cfg.duration_s) - track.shape[1]
            track = np.pad(track, ((0, 0), (0, pad_length)), mode="constant", constant_values=0)
        return track

    def loudness(self, audio_array):
        import pyloudnorm as pyln

        m = pyln.Meter(48000)  # create BS.1770 meter
        assert len(audio_array.shape) == 2
        lufs_db = m.integrated_loudness(audio_array.T)
        return lufs_db

    def normalize_volume(self, array_float, target_db=-16):
        loudness = self.loudness(array_float)
        
        # Handle edge cases where loudness measurement fails
        if not np.isfinite(loudness) or loudness < -80:
            # For very quiet/silent audio, return the original array
            # or apply minimal normalization
            max_val = np.abs(array_float).max()
            if max_val > 0:
                return array_float / max_val * 0.1  # Scale to 10% of max
            else:
                return array_float  # Return silent audio as-is
        
        gain_factor = np.log(10) / 20
        gain = target_db - loudness
        
        # Limit gain to reasonable bounds to prevent overflow
        gain = np.clip(gain, -40, 40)  # Limit to ±40dB adjustment
        
        gain = np.exp(gain * gain_factor)
        
        # Additional check to ensure gain is finite
        if not np.isfinite(gain):
            gain = 1.0
            
        norm_arr = array_float * gain
        
        # Prevent clipping
        if np.abs(norm_arr).max() > 1:
            norm_arr = norm_arr / np.abs(norm_arr).max()
            
        return norm_arr

    def clip_audio(self, array_float):
        if np.abs(array_float).max() > 1:
            print("Clipping audio by ", np.abs(array_float).max())
            return array_float / np.abs(array_float).max()
        else:
            return array_float

    def _load_audio(
        self,
        local_filepath,
        s3_filepath=None,
        expected_duration_s=180,
        audio_stats=None,
        max_duration_s=1,
        target_loudness_db=-16,
    ):
        try:
            start_s = (
                max(0.0, random.random() * int(expected_duration_s - self.audio_cfg.duration_s))
                if self.split == "train"
                else expected_duration_s / 2
            )
            audio = get_sample_oracle_file_segment(
                local_filepath=local_filepath,
                s3_filepath=s3_filepath,
                start_s=start_s,
                max_duration_s=max_duration_s,
            )
            assert audio.sample_rate == 48_000
            assert audio.n_channels == 2

            # normalize gain
            gain_db = 0.0
            if audio_stats is not None:
                loudness_db = audio_stats.get("loudness", None)
                if loudness_db is not None:
                    gain_db = target_loudness_db - float(loudness_db)
                    gain_db = np.clip(gain_db, -12, 12)
            audio, _ = audio.apply_gain(gain_db)

            target_sample_length = int(48000 * self.audio_cfg.duration_s)
            array_float = audio.array_float[:, :target_sample_length]
            array_float = self.pad_audio(array_float)
            array_float = self.clip_audio(array_float)

            return array_float, local_filepath, start_s
        except Exception as e:
            print(traceback.format_exc())
            print(
                f"Host {os.environ['HOSTNAME']} Error loading audio for {s3_filepath} "
                f"(local_fp: {local_filepath}, start_s: {start_s}, max_duration_s: {max_duration_s}): {e}"
            )

    def __next__(self):
        return self._next()

    def _sample_meta(self):
        """
        Randomly sample a meta from the metas list.
        Weighted sampling is slow for large weighted datasets so we sample 10k at a time.
        """
        if len(self.random_cache) == 0:
            choices = random.choices(range(len(self.metas)), k=10000)
            self.random_cache.extend(choices)
        idx = self.random_cache.pop()
        return self.metas[idx]

    def _next(self):
        # randomly sample a meta
        main_meta = self._sample_meta()

        # load wav
        data_wav, filepath, start_s = self._load_audio(
            main_meta["local_filepath"],
            main_meta.get("s3_filepath", None),
            main_meta["duration_s"],
            main_meta.get("audio_stats", None),
        )

        sample_data = SampleData(data_wav=data_wav, filepath=filepath, start_s=start_s)
        return sample_data
