from sunodata.dataset_maker_utils import DatasetConfig


class KaraokeStemsDataset(DatasetConfig):
    def _parse_arrays(self, meta_info, vae_arr):
        vae_start_idx = meta_info["offset_s"] * 25
        vae_end_idx = vae_start_idx + 25 * 30

        arr_v = vae_arr[vae_start_idx:vae_end_idx, :].copy()

        if meta_info.get("energy") is not None:
            energy_arr = meta_info["energy"]
            energy_arr = energy_arr[meta_info["offset_s"] : meta_info["offset_s"] + 30]
        else:
            energy_arr = None

        def expand_id(id):
            return f"{self.bundle.name}__{id}"

        if meta_info.get("stem_ids") is not None:
            stem_ids = meta_info["stem_ids"]
            stem_ids = [expand_id(id) for id in stem_ids]
        else:
            stem_ids = None

        new_meta = {
            "id": expand_id(meta_info["id"]),
            "bundle_id": meta_info["bundle_id"],
            "stem_type": meta_info["stem_type"] if "stem_type" in meta_info else None,
            "category": meta_info["category"] if "category" in meta_info else None,
            "type": meta_info["type"],
            "stem_ids": stem_ids,
            "energy": energy_arr,
            "duration_s": 30.0,
            "original_id": meta_info["id"],
            "original_offset_s": round(meta_info["offset_s"], 1),
            "original_s3_filepath": meta_info["s3_filepath"],
        }
        if "stem_type_id" in meta_info:
            new_meta["stem_type_id"] = meta_info["stem_type_id"]
        if "stem_type" in meta_info:
            new_meta["stem_type"] = meta_info["stem_type"]
        return [(arr_v, new_meta)]


class WetDryDataset(DatasetConfig):
    def _parse_arrays(self, meta_info, vae_arr):
        vae_start_idx = 0
        vae_end_idx = 25 * 30

        arr_v = vae_arr[vae_start_idx:vae_end_idx, :].copy()

        new_meta = meta_info.copy()
        new_meta["id"] = f"{self.bundle.name}__{meta_info['uuid']}"
        new_meta["bundle_id"] = meta_info["parent"]["uuid"]
        new_meta["tags"] = [t["label"] for t in meta_info["tags"]]
        new_meta["duration_s"] = 30.0
        new_meta["original_id"] = meta_info["uuid"]
        new_meta["original_offset_s"] = 0
        new_meta["s3_filepath"] = meta_info["s3_filepath"]
        new_meta["original_s3_filepath"] = meta_info["orig_s3_filepath"]

        if "stem_type_id" in meta_info:
            new_meta["stem_type_id"] = meta_info["stem_type_id"]
        if "stem_type" in meta_info:
            new_meta["stem_type"] = meta_info["stem_type"]
        return [(arr_v, new_meta)]


class RandomMixDataset(DatasetConfig):
    def _parse_arrays(self, meta_info, vae_arr):
        vae_start_idx = 0
        vae_end_idx = 25 * 30

        arr_v = vae_arr[vae_start_idx:vae_end_idx, :].copy()

        new_meta = meta_info.copy()
        new_meta["id"] = meta_info["id"]
        new_meta["bundle_id"] = meta_info["bundle_id"]
        new_meta["duration_s"] = 30.0
        new_meta["original_offset_s"] = 0

        return [(arr_v, new_meta)]
