import json
import boto3
import statistics
from typing import Dict, List, Any
import matplotlib.pyplot as plt
import numpy as np
import math
import pandas as pd
from tqdm import tqdm

from suno_utils.audio import Audio

def load_json_from_s3(bucket_name: str, key: str) -> Dict[str, Any]:
    """Load JSON data from S3 bucket"""
    s3 = boto3.client("s3")
    try:
        response = s3.get_object(Bucket=bucket_name, Key=key)
        content = response["Body"].read().decode("utf-8")
        return json.loads(content)
    except Exception as e:
        # print(f"Error loading {key}: {e}")
        return None


def process_wer_data(
    ids_data: Dict[
        str,
        List[str],
    ],
    bucket_name: str = "suno-data-uploads",
    s3_prefix: str = "tasks/feature_eval/cover_persona/2025_07_11-16_20_50/",
) -> Dict[str, Any]:
    """Process WER data for all files"""

    results = {}
    flat_data = []
    all_wers = []  # Collect all WER values for overall stats
    all_data = []  # Collect all data for overall stats

    for group_id, file_ids in tqdm(ids_data.items()):
        group_wers = []
        group_data = []

        for file_id in file_ids:
            # Construct S3 key
            s3_key = f"{s3_prefix}{file_id}_infill_wer.json"

            # Load JSON from S3
            data = load_json_from_s3(bucket_name, s3_key)

            if data and "wer" in data:
                data["s3_id"] = file_id
                flat_data.append(data)

                group_wers.append(data["wer"])
                group_data.append(data)

                # Add to overall collections
                all_wers.append(data["wer"])
                all_data.append(data)
            else:
                pass
                # print(f"Missing or invalid data for {file_id}")

        # Calculate statistics for this group
        if group_wers:
            results[group_id] = {
                "wers": group_wers,
                "mean_wer": statistics.mean(group_wers),
                "median_wer": statistics.median(group_wers),
                "min_wer": min(group_wers),
                "max_wer": max(group_wers),
                "std_wer": statistics.stdev(group_wers) if len(group_wers) > 1 else 0,
                "count": len(group_wers),
                "data": group_data,  # Include full data if needed
            }

    # Add overall statistics as "all" entry
    if all_wers:
        results["all"] = {
            "wers": all_wers,
            "mean_wer": statistics.mean(all_wers),
            "median_wer": statistics.median(all_wers),
            "min_wer": min(all_wers),
            "max_wer": max(all_wers),
            "std_wer": statistics.stdev(all_wers) if len(all_wers) > 1 else 0,
            "count": len(all_wers),
            "data": all_data,
        }

    return results, flat_data


def get_infill_data(timestamp):
    infill_path = f"/home/sara/glockenspiel/suno_utils/task_eval/modal_runs/infill_mappings_{timestamp}.json"
    with open(infill_path, "r") as f:
        ids_data = json.load(f)
    wer_results, flat_data = process_wer_data(
        ids_data, s3_prefix=f"tasks/feature_eval/cover_persona/{timestamp}/"
    )
    return wer_results, flat_data


def get_dur_success(flat_data):
    df = pd.DataFrame(flat_data)
    counts = df["duration_matches"].value_counts()
    print(f"{len(df)} items")
    percentages = counts / len(df)
    return percentages


def get_closest_bucket(infill_dur):
    buckets = [8, 15, 25]
    bucket_distance = {
        bucket_dur: abs(infill_dur - bucket_dur) for bucket_dur in buckets
    }
    min_key = min(bucket_distance, key=bucket_distance.get)
    return min_key


def get_wer_by_duration(flat_data):
    df = pd.DataFrame(flat_data)
    df["duration_bucket"] = df["infill_dur_s"].apply(get_closest_bucket)
    return df.groupby("duration_bucket")["wer"].mean()

timestamp = "2025_07_14-20_38_56"
results_30b, flat_30b = get_infill_data(timestamp)
print(get_dur_success(flat_30b))