#!/usr/bin/env python
from astropy.io import fits
import os
import sys
import argparse
from pathlib import Path
from multiprocessing import Pool, cpu_count
import time
import warnings

def split_det_file(input_file):
    """Split DET file into separate FITS files by extension type."""
    try:
        base_name = Path(input_file).stem
        
        with warnings.catch_warnings(record=True) as w:
            warnings.simplefilter("always")
            
            with fits.open(input_file, ignore_missing_simple=True) as hdulist:
                # Group extensions by type
                extensions = {'sci': [], 'flg': [], 'rms': []}
                
                for hdu in hdulist[1:]:  # Skip primary HDU
                    if hdu.name.endswith('.SCI'):
                        extensions['sci'].append(hdu)
                    elif hdu.name.endswith('.FLG'):
                        extensions['flg'].append(hdu)
                    elif hdu.name.endswith('.RMS'):
                        extensions['rms'].append(hdu)
                
                # Write separate files for each extension type that has data
                results = []
                for ext_type, hdus in extensions.items():
                    if hdus:
                        output_file = f"{base_name}_{ext_type}.fits"
                        new_hdulist = fits.HDUList([hdulist[0]] + hdus)
                        new_hdulist.writeto(output_file, overwrite=True)
                        results.append(f"{output_file} ({len(hdus)} extensions)")
                        print(f"Created: {output_file}")
                
                # Report any warnings
                if w:
                    for warning in w:
                        if 'truncated' in str(warning.message):
                            print(f"Warning for {input_file}: {warning.message}")
                
                return f"Processed {input_file}: " + ", ".join(results) if results else f"No .SCI/.FLG/.RMS extensions found in {input_file}"
    
    except Exception as e:
        error_msg = f"Error processing {input_file}: {e}"
        print(error_msg)
        return error_msg

def main():
    parser = argparse.ArgumentParser(description='Split DET FITS files into separate SCI/FLG/RMS extension files')
    parser.add_argument('files', nargs='+', help='Input FITS files to process')
    parser.add_argument('-j', '--jobs', type=int, default=min(4, cpu_count()), 
                       help=f'Number of parallel processes (default: {min(4, cpu_count())})')
    parser.add_argument('-v', '--verbose', action='store_true', 
                       help='Verbose output')
    parser.add_argument('--debug', action='store_true',
                       help='Show HDU structure (implies -j 1)')
    
    args = parser.parse_args()
    
    if args.debug:
        args.jobs = 1  # Force single-threaded for debug output
    
    # Validate input files
    valid_files = []
    for input_file in args.files:
        if not os.path.exists(input_file):
            print(f"Warning: {input_file} not found, skipping")
        else:
            valid_files.append(input_file)
    
    if not valid_files:
        print("No valid input files found")
        sys.exit(1)
    
    print(f"Processing {len(valid_files)} files using {args.jobs} processes...")
    
    start_time = time.time()
    
    if args.jobs == 1:
        # Single-threaded
        results = []
        for input_file in valid_files:
            if args.debug:
                # Show HDU structure for debugging
                with fits.open(input_file) as hdulist:
                    print(f"File {input_file} has {len(hdulist)} HDUs:")
                    for i, hdu in enumerate(hdulist[:10]):  # Show first 10
                        print(f"  HDU {i}: {hdu.name} ({type(hdu).__name__})")
                    if len(hdulist) > 10:
                        print(f"  ... and {len(hdulist)-10} more HDUs")
            
            result = split_det_file(input_file)
            results.append(result)
    else:
        # Multi-threaded processing
        with Pool(args.jobs) as pool:
            results = pool.map(split_det_file, valid_files)
    
    # Show results
    if args.verbose:
        print("\nSummary:")
        for result in results:
            print(result)
    
    elapsed = time.time() - start_time
    print(f"Completed processing {len(valid_files)} files in {elapsed:.2f} seconds")

if __name__ == "__main__":
    main()
