#!/usr/bin/env python3
"""
Wasabi Manifest Comparator
Processes partitions in small batches to avoid memory issues

Author: Gowtham Kamireddi
Co-Author: Francisco Santos

This optimized version:
- Faster partition splitting with batch writes
- Efficient file combining using pandas
- Reduced garbage collection calls
- Better chunking strategy (1-2M rows per chunk)
- Uses pandas groupby for partition filtering
- Consolidated glob operations

Usage:
    python compare_manifests.py
"""

import pandas as pd
import hashlib
import sys
from pathlib import Path
import time
from tqdm import tqdm
import json
import glob
import os
from typing import Dict, List, Tuple
import shutil

# Constants
DEFAULT_PARTITIONS = 100
CHUNK_SIZE = 1_500_000
PROGRESS_INTERVAL = 10
DEFAULT_OUTPUT_DIR = "./comparison_results/"
DEFAULT_PARTITIONS_DIR_A = "./partitions_source/"
DEFAULT_PARTITIONS_DIR_B = "./partitions_dest/"
DEFAULT_TEMP_DIR = "temp_results"
COLUMN_NAMES = ['bucket', 'file_path', 'object_id']
SEPARATOR_FULL = "=" * 80
SEPARATOR_DASHED = "─" * 80
PARTITION_FILE_PATTERN = "partition_{:03d}.csv"
DUP_PATTERN_A = "dup_a_{:03d}.csv"
DUP_PATTERN_B = "dup_b_{:03d}.csv"
MISS_PATTERN_B = "miss_b_{:03d}.csv"
MISS_PATTERN_A = "miss_a_{:03d}.csv"
DUP_GLOB_A = "dup_a_*.csv"
DUP_GLOB_B = "dup_b_*.csv"
MISS_GLOB_B = "miss_b_*.csv"
MISS_GLOB_A = "miss_a_*.csv"
SUMMARY_FILENAME = "comparison_summary.json"

class OptimizedManifestComparator:
    """
    Optimized comparator with faster I/O and reduced memory overhead
    """
    
    def __init__(self, config):
        self.config = config
        self.num_partitions = config.get('num_partitions', DEFAULT_PARTITIONS)
        self.label_a = config.get('label_a', 'FILE_A')
        self.label_b = config.get('label_b', 'FILE_B')
        self.output_dir = Path(config['output_dir'])
        self.partitions_dir_a = Path(config['partitions_dir_a'])
        self.partitions_dir_b = Path(config['partitions_dir_b'])
        
        # Create directories
        self.output_dir.mkdir(exist_ok=True, parents=True)
        self.partitions_dir_a.mkdir(exist_ok=True, parents=True)
        self.partitions_dir_b.mkdir(exist_ok=True, parents=True)
        
        self._print_header("Wasabi Manifest Comparator")
        print(f"{self.label_a}: {config['file_a']}")
        print(f"{self.label_b}: {config['file_b']}")
        print(f"Partitions: {self.num_partitions}")
        print(f"Mode: SEQUENTIAL (ultra-safe, no parallelization)")
        print(SEPARATOR_FULL)
    
    @staticmethod
    def _print_header(title: str) -> None:
        """Print formatted header"""
        print(SEPARATOR_FULL)
        print(title)
        print(SEPARATOR_FULL)
    
    @staticmethod
    def _validate_file_path(prompt: str, file_type: str = "file") -> str:
        """Validate and return an existing file path"""
        while True:
            file_path = input(prompt).strip().strip("'\"")
            if os.path.exists(file_path):
                return file_path
            print(f"❌ {file_type} not found: {file_path}")
            print("   Please check the path and try again\n")
    
    @staticmethod
    def hash_partition(file_path: str, num_partitions: int) -> int:
        """Generate partition number"""
        hash_val = int(hashlib.md5(file_path.encode()).hexdigest(), 16)
        return hash_val % num_partitions
    
    def split_manifest_file(self, input_file: str, output_dir: Path, file_label: str) -> None:
        """Split manifest into partitions - optimized batch writing
        
        Args:
            input_file: Path to the CSV manifest file to split
            output_dir: Directory to write partition files
            file_label: Label for this file (used in output messages)
        """
        print(f"\n{SEPARATOR_DASHED}")
        print(f"Splitting {file_label}...")
        print(SEPARATOR_DASHED)
        
        print("  📊 Counting rows...")
        total_rows = sum(1 for _ in open(input_file))
        print(f"  Total rows: {total_rows:,}")
        
        # Initialize partition data storage (dict of lists for batch writing)
        partition_data = {i: [] for i in range(self.num_partitions)}
        
        with tqdm(total=total_rows, desc=f"  Processing", unit=" rows") as pbar:
            for chunk in pd.read_csv(
                input_file, 
                chunksize=CHUNK_SIZE,
                names=COLUMN_NAMES,
                header=None
            ):
                # Add partition column
                chunk['_partition'] = chunk['file_path'].apply(
                    lambda path: self.hash_partition(path, self.num_partitions)
                )
                
                # Group by partition and collect data
                for partition_id, group in chunk.groupby('_partition'):
                    partition_data[partition_id].append(group[COLUMN_NAMES])
                
                pbar.update(len(chunk))
        
        # Write all partitions at once
        print("  Writing partitions...")
        for partition_id in range(self.num_partitions):
            if partition_data[partition_id]:
                combined = pd.concat(partition_data[partition_id], ignore_index=True)
                combined.to_csv(
                    output_dir / PARTITION_FILE_PATTERN.format(partition_id),
                    index=False,
                    header=False
                )
        
        print(f"  ✅ Created {self.num_partitions} partition files\n")
    
    def split_files(self) -> None:
        """Split both manifest files into partitions for parallel processing
        
        Executes PHASE 1 of comparison pipeline. Creates partition files
        in self.partitions_dir_a and self.partitions_dir_b.
        """
        self._print_header("PHASE 1: SPLITTING FILES")
        
        start_time = time.time()
        
        self.split_manifest_file(
            self.config['file_a'],
            self.partitions_dir_a,
            self.label_a
        )
        
        self.split_manifest_file(
            self.config['file_b'],
            self.partitions_dir_b,
            self.label_b
        )
        
        elapsed = time.time() - start_time
        print(f"✅ Splitting completed in {elapsed/60:.2f} minutes")
    
    def compare_partition_pair(self, partition_id: int, output_dir: Path) -> Dict:
        """Compare one partition pair and write results to disk
        
        Args:
            partition_id: Index of the partition to compare (0-99)
            output_dir: Directory to write comparison result files
            
        Returns:
            Dictionary containing comparison statistics for this partition
        """
        
        file_a = self.partitions_dir_a / PARTITION_FILE_PATTERN.format(partition_id)
        file_b = self.partitions_dir_b / PARTITION_FILE_PATTERN.format(partition_id)
        
        # Read data
        df_a = pd.read_csv(file_a, names=COLUMN_NAMES, header=None)
        df_b = pd.read_csv(file_b, names=COLUMN_NAMES, header=None)
        
        results = {
            'rows_a': len(df_a),
            'rows_b': len(df_b),
            'duplicates_a_count': 0,
            'duplicates_b_count': 0,
            'missing_in_b_count': 0,
            'missing_in_a_count': 0,
            'matching_count': 0
        }
        
        # 1. Duplicates
        duplicates_a = df_a[df_a.duplicated(subset=['file_path'], keep=False)]
        if not duplicates_a.empty:
            results['duplicates_a_count'] = len(duplicates_a)
            duplicates_a.to_csv(
                output_dir / DUP_PATTERN_A.format(partition_id),
                index=False,
                header=False
            )
        
        duplicates_b = df_b[df_b.duplicated(subset=['file_path'], keep=False)]
        if not duplicates_b.empty:
            results['duplicates_b_count'] = len(duplicates_b)
            duplicates_b.to_csv(
                output_dir / DUP_PATTERN_B.format(partition_id),
                index=False,
                header=False
            )
        
        # 2. Bidirectional comparison (using sets for efficiency)
        paths_a = set(df_a['file_path'])
        paths_b = set(df_b['file_path'])
        
        # Missing in B
        missing_paths_b = paths_a - paths_b
        if missing_paths_b:
            missing_in_b = df_a[df_a['file_path'].isin(missing_paths_b)]
            results['missing_in_b_count'] = len(missing_in_b)
            missing_in_b.to_csv(
                output_dir / MISS_PATTERN_B.format(partition_id),
                index=False,
                header=False
            )
        
        # Missing in A
        missing_paths_a = paths_b - paths_a
        if missing_paths_a:
            missing_in_a = df_b[df_b['file_path'].isin(missing_paths_a)]
            results['missing_in_a_count'] = len(missing_in_a)
            missing_in_a.to_csv(
                output_dir / MISS_PATTERN_A.format(partition_id),
                index=False,
                header=False
            )
        
        # Matching
        results['matching_count'] = len(paths_a & paths_b)
        
        return results
    
    def compare_all_partitions(self) -> Tuple[List[Dict], Path]:
        """Compare all partitions SEQUENTIALLY (no multiprocessing)"""
        self._print_header("PHASE 2: SEQUENTIAL COMPARISON (NO PARALLELIZATION)")
        print("⚠️  This will be slower but won't crash due to memory\n")
        
        output_dir = self.output_dir / DEFAULT_TEMP_DIR
        output_dir.mkdir(exist_ok=True)
        
        start_time = time.time()
        results = []
        
        for i in tqdm(range(self.num_partitions), desc="  Comparing", unit=" partition"):
            result = self.compare_partition_pair(i, output_dir)
            results.append(result)
            
            # Progress update every N partitions
            if (i + 1) % PROGRESS_INTERVAL == 0:
                elapsed = time.time() - start_time
                avg_time = elapsed / (i + 1)
                remaining = avg_time * (self.num_partitions - i - 1)
                print(f"    Progress: {i+1}/{self.num_partitions} | "
                      f"Elapsed: {elapsed/60:.1f}m | "
                      f"ETA: {remaining/60:.1f}m")
        
        elapsed = time.time() - start_time
        print(f"\n✅ Comparison completed in {elapsed/60:.2f} minutes")
        
        return results, output_dir
    
    def aggregate_results(self, partition_results: List[Dict], temp_dir: Path) -> Dict:
        """Aggregate results from all partitions and generate final report
        
        Args:
            partition_results: List of result dictionaries from each partition
            temp_dir: Temporary directory containing partition comparison files
            
        Returns:
            Dictionary containing aggregated summary statistics
        """
        self._print_header("PHASE 3: AGGREGATING RESULTS")
        
        # Calculate summary using vectorized operations
        summary = {
            'total_rows_a': sum(r['rows_a'] for r in partition_results),
            'total_rows_b': sum(r['rows_b'] for r in partition_results),
            'total_duplicates_a': sum(r['duplicates_a_count'] for r in partition_results),
            'total_duplicates_b': sum(r['duplicates_b_count'] for r in partition_results),
            'total_missing_in_b': sum(r['missing_in_b_count'] for r in partition_results),
            'total_missing_in_a': sum(r['missing_in_a_count'] for r in partition_results),
            'total_matching': sum(r['matching_count'] for r in partition_results),
            'label_a': self.label_a,
            'label_b': self.label_b
        }
        
        # Save summary
        with open(self.output_dir / SUMMARY_FILENAME, 'w') as f:
            json.dump(summary, f, indent=2)
        print("  ✅ Summary saved")
        
        # Combine files efficiently
        for pattern, output_name, label in self._get_file_pairs_for_aggregation():
            files = sorted(glob.glob(str(temp_dir / pattern)))
            if files:
                dfs = [pd.read_csv(f, names=COLUMN_NAMES, header=None) for f in files]
                combined = pd.concat(dfs, ignore_index=True)
                combined.to_csv(
                    self.output_dir / output_name,
                    index=False,
                    header=False
                )
                print(f"  ✅ {label}: {len(combined):,} rows")
        
        return summary
    
    def _get_file_pairs_for_aggregation(self) -> List[Tuple[str, str, str]]:
        """Build list of file patterns and output names for aggregation
        
        Returns:
            List of tuples: (glob_pattern, output_filename, description_label)
        """
        return [
            (DUP_GLOB_A, f'duplicates_in_{self.label_a}.csv', f'Duplicates in {self.label_a}'),
            (DUP_GLOB_B, f'duplicates_in_{self.label_b}.csv', f'Duplicates in {self.label_b}'),
            (MISS_GLOB_B, f'Found_in_{self.label_a}_but_NOT_in_{self.label_b}.csv', f'Found in {self.label_a} but not in the {self.label_b}'),
            (MISS_GLOB_A, f'Found_in_{self.label_b}_but_NOT_in_{self.label_a}.csv', f'Found in {self.label_b} but not in the {self.label_a}'),
        ]
    
    def cleanup(self, temp_dir: Path) -> None:
        """Cleanup temporary partition and result files
        
        Args:
            temp_dir: Temporary directory to remove
        """
        print("\n🧹 Cleaning up...")
        
        if temp_dir.exists():
            shutil.rmtree(temp_dir)
        
        if self.config.get('cleanup_partitions', True):
            shutil.rmtree(self.partitions_dir_a, ignore_errors=True)
            shutil.rmtree(self.partitions_dir_b, ignore_errors=True)
        
        print("  ✅ Cleanup complete")
    
    def run(self) -> Dict:
        """Run full comparison"""
        self._print_header("🚀 STARTING COMPARISON")
        
        overall_start = time.time()
        
        # Phase 1: Split
        self.split_files()
        
        # Phase 2: Compare
        partition_results, temp_dir = self.compare_all_partitions()
        
        # Phase 3: Aggregate
        summary = self.aggregate_results(partition_results, temp_dir)
        
        # Phase 4: Cleanup
        self.cleanup(temp_dir)
        
        # Print summary
        total_time = time.time() - overall_start
        self._print_header("✅ COMPARISON COMPLETE!")
        print(f"⏱️  Total time: {total_time/60:.2f} minutes ({total_time/3600:.2f} hours)")
        print("\n📊 RESULTS:")
        print(SEPARATOR_DASHED)
        print(f"{self.label_a}: {summary['total_rows_a']:,} rows")
        print(f"{self.label_b}: {summary['total_rows_b']:,} rows")
        print(f"Matching: {summary['total_matching']:,} ✅")
        
        missing_b = summary['total_missing_in_b']
        missing_a = summary['total_missing_in_a']
        print(f"Missing in {self.label_b}: {missing_b:,} {'⚠️' if missing_b > 0 else '✅'}")
        print(f"Missing in {self.label_a}: {missing_a:,} {'⚠️' if missing_a > 0 else '✅'}")
        
        if summary['total_rows_a'] > 0:
            pct = (summary['total_matching'] / summary['total_rows_a']) * 100
            print(f"\n📈 Replication coverage: {pct:.2f}%")
        
        print(SEPARATOR_DASHED)
        print(f"📁 Results: {self.output_dir}")
        print(SEPARATOR_FULL)
        
        return summary


if __name__ == '__main__':
    
    print("\n" + SEPARATOR_FULL)
    print("MANIFEST COMPARATOR - INTERACTIVE MODE")
    print(SEPARATOR_FULL)
    
    # Prompt for file paths
    print("\n📁 Please provide the manifest file paths:")
    print("   (You can drag and drop files into the terminal)\n")
    
    source_file = OptimizedManifestComparator._validate_file_path("Source manifest file: ")
    dest_file = OptimizedManifestComparator._validate_file_path("Destination manifest file: ")
    
    # Optional: custom labels
    print("\n🏷️  Labels for output files (press Enter for defaults):")
    label_source = input("Source label [SOURCE]: ").strip() or "SOURCE"
    label_dest = input("Destination label [DESTINATION]: ").strip() or "DESTINATION"
    
    # Optional: output directory
    print("\n📂 Output directory (press Enter for default):")
    output_dir = input(f"Output directory [{DEFAULT_OUTPUT_DIR}]: ").strip() or DEFAULT_OUTPUT_DIR
    
    # Configuration
    config = {
        'file_a': source_file,
        'file_b': dest_file,
        'label_a': label_source,
        'label_b': label_dest,
        'output_dir': output_dir,
        'partitions_dir_a': DEFAULT_PARTITIONS_DIR_A,
        'partitions_dir_b': DEFAULT_PARTITIONS_DIR_B,
        'num_partitions': DEFAULT_PARTITIONS,
        'cleanup_partitions': True
    }
    
    print("\n⚠️  PERFORMANCE NOTE:")
    print("   This OPTIMIZED version uses faster I/O and better chunking")
    print("   Expected time: 30-60 minutes (3-4x faster than original!)")
    print("   Sequential processing keeps memory usage under control\n")
    
    print(SEPARATOR_FULL)
    print("CONFIGURATION SUMMARY:")
    print(SEPARATOR_FULL)
    print(f"Source: {source_file}")
    print(f"Destination: {dest_file}")
    print(f"Source label: {label_source}")
    print(f"Destination label: {label_dest}")
    print(f"Output directory: {output_dir}")
    print(f"Partitions: {DEFAULT_PARTITIONS}")
    print(SEPARATOR_FULL)
    
    input("\nPress Enter to start comparison, or Ctrl+C to cancel...")
    
    try:
        comparator = OptimizedManifestComparator(config)
        results = comparator.run()
        
        print("\n✅ SUCCESS! Check the comparison_results/ directory")
        
    except KeyboardInterrupt:
        print("\n\n⚠️  Interrupted by user")
        sys.exit(1)
    except Exception as e:
        print(f"\n\n❌ ERROR: {e}")
        import traceback
        traceback.print_exc()
