from suno_utils.worker.settings import s3_client
from pathlib import Path
from concurrent.futures import ThreadPoolExecutor, as_completed
from tqdm import tqdm
from zipfile import ZipFile, BadZipFile
import numpy as np
import io
import json
import random
import bisect


def _write_beats(f, data):
    downbeats = list(sorted(data["downbeat_times"].tolist()))
    beats = list(sorted(data["beat_times"].tolist()))
    bars = []

    for left_downbeat, right_downbeat in zip(
        [-float("inf")] + downbeats, downbeats + [float("inf")]
    ):
        start_idx = bisect.bisect_left(beats, left_downbeat)
        end_idx = bisect.bisect_left(beats, right_downbeat)
        beats_in_bar = end_idx - start_idx
        if beats_in_bar == 0:
            continue
        bars.append(beats[start_idx:end_idx])

    result = []
    rest_idx = 0
    if len(bars) > 1:
        # fix up the first bar
        offset = len(bars[1]) - len(bars[0])
        if offset < 0:
            offset = 0
        result.extend([b, i + 1 + offset] for i, b in enumerate(bars[0]))
        rest_idx = 1

    for bar in bars[rest_idx:]:
        result.extend([b, i + 1] for i, b in enumerate(bar))

    for b, d in result:
        f.write(f"{b}  {d}\n")


def main(suffix: str):
    paginator = s3_client.get_paginator("list_objects_v2")
    pages = paginator.paginate(
        Bucket="suno-data",
        Prefix=f"m4burns/downbeats_processed/spectrograms_round5_{suffix}/",
    )
    data_dir = Path(__file__).parent.parent / "data"
    # temp_dir = data_dir / "suno_download"
    temp_dir = Path(f"/mnt/localdisk/tmp_m4burns/beat_this/{suffix}")
    temp_dir.mkdir(parents=True, exist_ok=True)
    spect_uuids = []

    def download_file(obj):
        if not obj["Key"].endswith(".npz"):
            return None
        npz_name = obj["Key"].split("/")[-1]
        npz_out_path = temp_dir / npz_name
        spect_uuids.append(npz_name.split(".")[0])
        if npz_out_path.exists():
            return None
        try:
            with open(npz_out_path, "wb") as f:
                s3_client.download_fileobj(
                    Bucket="suno-data", Key=obj["Key"], Fileobj=f
                )
        except Exception:
            npz_out_path.unlink()
            raise

        return npz_name

    with ThreadPoolExecutor(max_workers=16) as executor:
        futures = [
            executor.submit(download_file, obj)
            for page in pages
            for obj in page.get("Contents", [])
        ]
        for future in tqdm(
            as_completed(futures), total=len(futures), desc="Downloading files"
        ):
            future.result()

    suno_ann_dir = data_dir / "annotations" / f"suno_synth_{suffix}"
    suno_ann_dir.mkdir(parents=True, exist_ok=True)
    with open(suno_ann_dir / "info.json", "w") as f:
        json.dump({"has_downbeats": True}, f)
    with open(suno_ann_dir / "single.split", "w") as f:
        for u in spect_uuids:
            split = "train" if random.random() < 0.9 else "val"
            f.write(f"{u}\t{split}\n")

    beats_dir = suno_ann_dir / "annotations" / "beats"
    beats_dir.mkdir(parents=True, exist_ok=True)

    with ZipFile(
        data_dir / "audio" / "spectrograms" / f"suno_synth_{suffix}.npz", "w"
    ) as z:
        for u in tqdm(spect_uuids, desc="Writing bundle+annotations"):
            try:
                data = np.load(temp_dir / f"{u}.npz")
            except (BadZipFile, EOFError):
                print(f"Bad zip file: {u}")
                continue
            buf = io.BytesIO()
            spec_array = data["spec_array"]
            if spec_array.ndim == 3 and spec_array.shape[-1] == 1:
                spec_array = spec_array.squeeze(-1)
            np.save(buf, spec_array.astype(np.float16))
            z.writestr(f"{u}/track.npy", buf.getvalue())
            with open(beats_dir / f"{u}.beats", "w") as f:
                _write_beats(f, data)


if __name__ == "__main__":
    import sys

    assert len(sys.argv) == 2, "Usage: python download_suno.py <suffix>"
    main(sys.argv[1])
