#!/usr/bin/env python3
"""
Parallel audio conversion to opus format.
Processes audio files from metadata, converting non-opus files to opus format.
"""

import os
import argparse
from concurrent.futures import ProcessPoolExecutor, as_completed
from functools import partial
from typing import Dict, Optional

from suno_utils.utils.text import read_jsonl
from suno_utils.audio import Audio
from tqdm import tqdm


def process_audio_item(meta_item: Dict, output_dir: str, skip_existing: bool = True) -> tuple[str, str, Optional[str]]:
    """
    Process a single audio item, ensuring it exists as opus in output directory.
    
    Args:
        meta_item: Metadata dictionary with 'id' and 's3_filepath' keys
        output_dir: Directory to save opus files
        skip_existing: If True, skip files that already exist in output directory
        
    Returns:
        Tuple of (item_id, status, error_message)
        status can be: 'converted', 'already_exists', 'copied_opus', 'error'
    """
    item_id = meta_item["id"]
    s3_path = meta_item["s3_filepath"]
    output_path = os.path.join(output_dir, f"{item_id}.opus")
    
    try:
        # Check if output file already exists
        if skip_existing and os.path.exists(output_path):
            return (item_id, "already_exists", None)
        
        # Load audio from S3
        audio = Audio.from_s3(s3_path)
        
        # Write to output directory (converts if needed, copies if already opus)
        audio.write_opus(output_path)
        
        # Determine if it was a conversion or just a copy
        if s3_path.endswith(".opus"):
            return (item_id, "copied_opus", None)
        else:
            return (item_id, "converted", None)
        
    except Exception as e:
        return (item_id, "error", str(e))


def main():
    parser = argparse.ArgumentParser(
        description="Consolidate audio files to opus format in parallel. "
                    "Converts non-opus files and copies opus files to output directory."
    )
    parser.add_argument(
        "--input",
        type=str,
        default="/app2/suno/data/sara/sfx_get_beats/combined_v3_w_extreme_metas_v0.jsonl",
        help="Input JSONL file with metadata"
    )
    parser.add_argument(
        "--output-dir",
        type=str,
        default="/app2/suno/data/sara/sfx_get_beats/opus_audio",
        help="Output directory for opus files"
    )
    parser.add_argument(
        "--workers",
        type=int,
        default=None,
        help="Number of parallel workers (default: CPU count)"
    )
    parser.add_argument(
        "--limit",
        type=int,
        default=None,
        help="Limit number of files to process (for testing)"
    )
    parser.add_argument(
        "--error-log",
        type=str,
        default="conversion_errors.log",
        help="File to log errors"
    )
    parser.add_argument(
        "--skip-existing",
        action="store_true",
        default=True,
        help="Skip files that already exist in output directory (default: True)"
    )
    parser.add_argument(
        "--no-skip-existing",
        dest="skip_existing",
        action="store_false",
        help="Reprocess all files, overwriting existing ones"
    )
    
    args = parser.parse_args()
    
    # Load metadata
    print(f"Loading metadata from {args.input}...")
    meta = read_jsonl(args.input)
    print(f"Loaded {len(meta)} items")
    
    # Apply limit if specified
    if args.limit:
        meta = meta[:args.limit]
        print(f"Processing limited to {len(meta)} items")
    
    # Ensure output directory exists
    os.makedirs(args.output_dir, exist_ok=True)
    
    # Check how many files already exist if skipping
    if args.skip_existing:
        existing_count = sum(1 for item in meta if os.path.exists(
            os.path.join(args.output_dir, f"{item['id']}.opus")
        ))
        print(f"Found {existing_count} existing files that will be skipped")
    
    # Process in parallel
    print(f"Starting parallel processing with {args.workers or 'default'} workers...")
    print(f"Skip existing: {args.skip_existing}")
    
    converted_count = 0
    already_exists_count = 0
    copied_opus_count = 0
    error_count = 0
    
    # Open error log file
    with open(args.error_log, 'w') as error_file:
        with ProcessPoolExecutor(max_workers=args.workers) as executor:
            # Submit all jobs
            process_func = partial(process_audio_item, output_dir=args.output_dir, 
                                  skip_existing=args.skip_existing)
            futures = {
                executor.submit(process_func, item): item 
                for item in meta
            }
            
            # Process completed jobs with progress bar
            with tqdm(total=len(futures), desc="Converting to opus", 
                     unit="files", ncols=120, colour='green') as pbar:
                for future in as_completed(futures):
                    item_id, status, error_msg = future.result()
                    
                    if status == "converted":
                        converted_count += 1
                    elif status == "already_exists":
                        already_exists_count += 1
                    elif status == "copied_opus":
                        copied_opus_count += 1
                    elif status == "error":
                        error_count += 1
                        error_file.write(f"{item_id}: {error_msg}\n")
                        error_file.flush()
                    
                    # Update progress bar with stats
                    pbar.set_postfix({
                        'converted': converted_count,
                        'copied': copied_opus_count,
                        'exists': already_exists_count,
                        'errors': error_count
                    }, refresh=True)
                    pbar.update(1)
    
    # Print summary
    print("\n" + "="*60)
    print("CONVERSION COMPLETE")
    print("="*60)
    print(f"Total processed:        {len(meta)}")
    print(f"Converted to opus:      {converted_count}")
    print(f"Copied opus files:      {copied_opus_count}")
    print(f"Already existed:        {already_exists_count}")
    print(f"Errors:                 {error_count}")
    
    total_new = converted_count + copied_opus_count
    print(f"\nTotal new files added:  {total_new}")
    
    if error_count > 0:
        print(f"\nErrors logged to: {args.error_log}")
    
    if args.skip_existing and already_exists_count > 0:
        print(f"\nℹ️  Skipped {already_exists_count} existing files. Use --no-skip-existing to reprocess them.")


if __name__ == "__main__":
    main()

