# %%
import json
from suno_utils.audio import Audio
import os
from openai import OpenAI

with open("/app2/suno/data/podcast/podcast_episodes_with_language.jsonl", "r") as f:
    metas = [json.loads(line) for line in f]

print(len(metas))

metas[0]


# %%
import matplotlib.pyplot as plt
from collections import Counter

languages = [meta["detected_language"] for meta in metas]
print(len(languages))

# Count language frequencies
lang_counts = Counter(languages)
print(f"Number of unique languages: {len(lang_counts)}")

# Create a bar plot
plt.figure(figsize=(12, 6))
langs, counts = zip(*lang_counts.most_common(20))  # Top 20 languages
plt.bar(langs, counts)
plt.xlabel("Language")
plt.ylabel("Count")
plt.title("Distribution of Detected Languages in Podcast Episodes")
plt.xticks(rotation=45, ha="right")
plt.tight_layout()
plt.show()

# %%
lang_counts

# %%
from collections import defaultdict
import random

# Group metas by language
metas_by_lang = defaultdict(list)
for meta in metas:
    metas_by_lang[meta["detected_language"]].append(meta)

# Stratify to keep up to 20k per language
stratified_metas = []
for lang, lang_metas in metas_by_lang.items():
    N = 1000
    if len(lang_metas) > N:
        # Randomly sample 20k if more than 20k
        sampled = random.sample(lang_metas, N)
        stratified_metas.extend(sampled)
    else:
        # Keep all if 20k or fewer
        stratified_metas.extend(lang_metas)

print(f"Original count: {len(metas)}")
print(f"Stratified count: {len(stratified_metas)}")

# Update metas to the stratified version
metas = stratified_metas


# %% cell 4.5 code

# Calculate total duration
total_duration_hours = sum(meta["duration_s"] for meta in metas) / 3600
print(f"Total duration: {total_duration_hours:.2f} hours")


# %%
metas_by_lang["english"][0]

path = f"/app2/suno/data/podcast/audio/{metas_by_lang['english'][0]['id']}.mp3"
audio = Audio.from_file(path)
audio.get_segment(0, 60 * 5).play()

# %%

# %%

# %%

# %%
audio.duration_s

# %%
from pydub import AudioSegment
from pydub.silence import split_on_silence


def split_audio_into_chunks(audio_path):
    """Split audio into roughly 5-minute chunks using silence detection."""
    # Load audio using pydub
    audio = AudioSegment.from_file(audio_path)

    # Split on silence to get natural boundaries
    chunks = split_on_silence(
        audio,
        min_silence_len=1000,  # minimum silence length in ms
        keep_silence=100,
        seek_step=300,
        silence_thresh=-40,
    )
    # print(f"Split audio into {len(chunks)} chunks")
    return chunks


# Test the function
test_audio_path = f"/app2/suno/data/podcast/audio/{metas_by_lang['english'][0]['id']}.mp3"
audio_chunks = split_audio_into_chunks(test_audio_path)
print(f"Split audio into {len(audio_chunks)} chunks")
for i, chunk in enumerate(audio_chunks):
    print(f"Chunk {i}: duration {len(chunk) / 1000:.1f} seconds")


# %%
from tqdm import tqdm
import concurrent.futures
import threading

api_key = os.getenv("OPENAI_API_KEY", "sk-proj-b6DmKnBFU0ELNITxlVoBT3BlbkFJEwSAkwXgzd82EcMsLE5Z")
if not api_key:
    raise ValueError("OPENAI_API_KEY environment variable is required")

client = OpenAI(api_key=api_key)


def transcribe_chunk(chunk_data):
    """Transcribe a single chunk."""
    i, chunk, path = chunk_data
    chunk_file = f"/app2/suno/data/podcast/audio_chunks/chunk_{i}_{path.split('/')[-1]}.mp3"
    chunk.export(chunk_file, format="mp3")

    # Transcribe the chunk
    with open(chunk_file, "rb") as f:
        transcript = client.audio.transcriptions.create(
            file=f,
            model="gpt-4o-mini-transcribe",
            # model="whisper-1",
            response_format="text",
        )
        # print(transcript)

    return {
        "chunk_index": i,
        "text": transcript,
        "duration": len(chunk) / 1000,  # duration in seconds
        "audio_path": chunk_file,
    }


def transcribe(path):
    """Split audio into chunks at natural silence boundaries and transcribe each chunk."""
    # Split audio into chunks at silence boundaries
    chunks = split_audio_into_chunks(path)
    # print(f"Split audio into {len(chunks)} chunks")
    # filter out chunks that are too long
    chunks = [chunk for chunk in chunks if 5_000 < len(chunk) < 8 * 60 * 1000]

    # Create output directory if it doesn't exist
    os.makedirs("/app2/suno/data/podcast/audio_chunks", exist_ok=True)
    # Prepare chunk data for concurrent processing
    chunk_data = [(i, chunk, path) for i, chunk in enumerate(chunks)]
    # Process chunks concurrently with threads
    with concurrent.futures.ThreadPoolExecutor(max_workers=10) as executor:
        transcripts = list(executor.map(transcribe_chunk, chunk_data))
    # for chunk in tqdm(chunk_data, disable=True):
    #     transcripts.append(transcribe_chunk(chunk))

    # Sort by chunk index to maintain order
    transcripts.sort(key=lambda x: x["chunk_index"])

    return transcripts


# transcribe(path)


# %%
def process_meta(meta):
    """Process a single metadata entry for transcription."""
    if "chunks" not in meta:
        try:
            transcript = transcribe(
                meta["s3_filepath"].replace(
                    "s3://suno-data/datasets/harvest/podcast_episodes/audio/",
                    "/app2/suno/data/podcast/audio/",
                )
            )
            meta["chunks"] = transcript
        except Exception as e:
            print(f"Error transcribing {meta['id']}: {e}")
            meta["chunks"] = None
    return meta


# Use ThreadPoolExecutor for concurrent processing with timeout
def process_with_timeout(meta):
    """Wrapper to process meta with timeout."""
    with concurrent.futures.ThreadPoolExecutor(max_workers=1) as executor:
        future = executor.submit(process_meta, meta)
        try:
            return future.result(timeout=300)  # 5 minute timeout per item
        except concurrent.futures.TimeoutError:
            print(f"Timeout processing {meta['id']}")
            meta["chunks"] = None
            return meta


# metas = random.sample(metas, 10)
with concurrent.futures.ThreadPoolExecutor(max_workers=10) as executor:
    metas = list(tqdm(executor.map(process_with_timeout, metas), total=len(metas)))

# Write updated metadata as JSONL
with open("/app2/suno/data/podcast/stratified_metadata_v2.jsonl", "w") as f:
    for meta in metas:
        f.write(json.dumps(meta) + "\n")
