#!/usr/bin/env python3
"""
Efficiently compare two metadata JSONL files with identical line counts.

Uses parallel processing to quickly identify any differences between files.
"""

import argparse
import json
import multiprocessing as mp
from pathlib import Path
from typing import Dict, List, Optional, Tuple

from tqdm import tqdm


def count_lines(file_path: Path) -> int:
    """Count lines in a file efficiently."""
    count = 0
    with open(file_path, "rb") as f:
        for _ in f:
            count += 1
    return count


def compare_records(line1: str, line2: str, line_num: int) -> Optional[Tuple[int, Dict]]:
    """
    Compare two JSONL records.

    Args:
        line1: First record as string
        line2: Second record as string
        line_num: Line number (1-indexed)

    Returns:
        None if records are identical, otherwise (line_num, diff_info)
    """
    try:
        record1 = json.loads(line1.strip())
        record2 = json.loads(line2.strip())

        if record1 == record2:
            return None

        # Records differ - find what's different
        diff_info = {
            "line_number": line_num,
            "id1": record1.get("id", "MISSING"),
            "id2": record2.get("id", "MISSING"),
        }

        # Find differing keys
        keys1 = set(record1.keys())
        keys2 = set(record2.keys())

        only_in_1 = keys1 - keys2
        only_in_2 = keys2 - keys1
        common_keys = keys1 & keys2

        if only_in_1:
            diff_info["only_in_file1"] = list(only_in_1)
        if only_in_2:
            diff_info["only_in_file2"] = list(only_in_2)

        # Check common keys for value differences
        differing_keys = []
        for key in common_keys:
            if record1[key] != record2[key]:
                differing_keys.append(
                    {
                        "key": key,
                        "value1": record1[key]
                        if not isinstance(record1[key], (list, dict))
                        else f"<{type(record1[key]).__name__}>",
                        "value2": record2[key]
                        if not isinstance(record2[key], (list, dict))
                        else f"<{type(record2[key]).__name__}>",
                    }
                )

        if differing_keys:
            diff_info["differing_values"] = differing_keys

        return line_num, diff_info

    except json.JSONDecodeError as e:
        return line_num, {
            "line_number": line_num,
            "error": "JSON parse error",
            "details": str(e),
        }
    except Exception as e:
        return line_num, {
            "line_number": line_num,
            "error": "Comparison error",
            "details": str(e),
        }


def process_chunk(args: Tuple[List[str], List[str], int]) -> List[Tuple[int, Dict]]:
    """
    Process a chunk of lines in parallel.

    Args:
        args: Tuple of (lines1, lines2, start_line_num)

    Returns:
        List of differences found
    """
    lines1, lines2, start_line_num = args
    differences = []

    for i, (line1, line2) in enumerate(zip(lines1, lines2)):
        line_num = start_line_num + i
        result = compare_records(line1, line2, line_num)
        if result is not None:
            differences.append(result)

    return differences


def compare_files_parallel(
    file1: Path,
    file2: Path,
    chunk_size: int = 10000,
    num_workers: Optional[int] = None,
) -> List[Tuple[int, Dict]]:
    """
    Compare two JSONL files in parallel.

    Args:
        file1: First file path
        file2: Second file path
        chunk_size: Number of lines to process per chunk
        num_workers: Number of worker processes (defaults to CPU count)

    Returns:
        List of differences found (line_num, diff_info)
    """
    if num_workers is None:
        num_workers = mp.cpu_count()

    print(f"Using {num_workers} worker processes with chunk size {chunk_size}")

    # Get total lines
    print("Counting lines...")
    total_lines = count_lines(file1)
    print(f"Total lines: {total_lines:,}")

    # Verify both files have same line count
    lines2 = count_lines(file2)
    if total_lines != lines2:
        raise ValueError(f"Files have different line counts: {total_lines} vs {lines2}")

    differences = []

    # Process files in chunks
    with mp.Pool(processes=num_workers) as pool:
        with open(file1, "r") as f1, open(file2, "r") as f2:
            chunk_args = []
            line_num = 1

            # Read and prepare chunks
            print("Processing files in parallel...")
            with tqdm(total=total_lines, unit="lines") as pbar:
                while True:
                    # Read chunk from both files
                    lines1 = []
                    lines2 = []

                    for _ in range(chunk_size):
                        line1 = f1.readline()
                        line2 = f2.readline()

                        if not line1:  # EOF
                            break

                        lines1.append(line1)
                        lines2.append(line2)

                    if not lines1:  # No more lines
                        break

                    chunk_args.append((lines1, lines2, line_num))
                    line_num += len(lines1)

                    # Process chunks when we have enough or reached EOF
                    if len(chunk_args) >= num_workers or len(lines1) < chunk_size:
                        # Process accumulated chunks
                        results = pool.map(process_chunk, chunk_args)

                        # Collect differences
                        for chunk_diffs in results:
                            differences.extend(chunk_diffs)

                        # Update progress
                        processed = sum(len(args[0]) for args in chunk_args)
                        pbar.update(processed)

                        # Reset for next batch
                        chunk_args = []

    return differences


def main():
    parser = argparse.ArgumentParser(
        description="Compare two metadata JSONL files with identical line counts"
    )

    parser.add_argument("file1", type=Path, help="First JSONL file path")
    parser.add_argument("file2", type=Path, help="Second JSONL file path")

    parser.add_argument(
        "--chunk-size",
        type=int,
        default=10000,
        help="Number of lines to process per chunk (default: 10000)",
    )

    parser.add_argument(
        "--workers",
        type=int,
        default=None,
        help="Number of worker processes (default: CPU count)",
    )

    parser.add_argument(
        "--max-diffs",
        type=int,
        default=100,
        help="Maximum number of differences to display (default: 100)",
    )

    parser.add_argument(
        "--output",
        type=Path,
        default=None,
        help="Save detailed differences to JSON file",
    )

    args = parser.parse_args()

    # Validate files
    if not args.file1.exists():
        print(f"Error: File not found: {args.file1}")
        return 1

    if not args.file2.exists():
        print(f"Error: File not found: {args.file2}")
        return 1

    print("\n" + "=" * 80)
    print("METADATA FILE COMPARISON")
    print("=" * 80)
    print(f"File 1: {args.file1}")
    print(f"File 2: {args.file2}")
    print("=" * 80)

    try:
        # Compare files
        differences = compare_files_parallel(
            args.file1, args.file2, chunk_size=args.chunk_size, num_workers=args.workers
        )

        # Report results
        print("\n" + "=" * 80)
        print("RESULTS")
        print("=" * 80)

        if not differences:
            print("✓ Files are IDENTICAL - no differences found!")
        else:
            print(f"✗ Files DIFFER - found {len(differences):,} differences")

            # Display first N differences
            print(f"\nShowing first {min(len(differences), args.max_diffs)} differences:")
            print("-" * 80)

            for i, (line_num, diff_info) in enumerate(differences[: args.max_diffs], 1):
                print(f"\nDifference #{i} at line {line_num}:")

                if "error" in diff_info:
                    print(f"  ERROR: {diff_info['error']}")
                    print(f"  Details: {diff_info['details']}")
                else:
                    print(f"  ID in file1: {diff_info.get('id1', 'N/A')}")
                    print(f"  ID in file2: {diff_info.get('id2', 'N/A')}")

                    if "only_in_file1" in diff_info:
                        print(f"  Keys only in file1: {diff_info['only_in_file1']}")

                    if "only_in_file2" in diff_info:
                        print(f"  Keys only in file2: {diff_info['only_in_file2']}")

                    if "differing_values" in diff_info:
                        print("  Differing values:")
                        for kv in diff_info["differing_values"][:5]:  # Show first 5
                            print(f"    {kv['key']}:")
                            print(f"      file1: {kv['value1']}")
                            print(f"      file2: {kv['value2']}")

                        if len(diff_info["differing_values"]) > 5:
                            remaining = len(diff_info["differing_values"]) - 5
                            print(f"    ... and {remaining} more differing keys")

            if len(differences) > args.max_diffs:
                print(f"\n... and {len(differences) - args.max_diffs} more differences")

        # Save detailed output if requested
        if args.output:
            output_data = {
                "file1": str(args.file1),
                "file2": str(args.file2),
                "total_lines_compared": count_lines(args.file1),
                "total_differences": len(differences),
                "differences": [
                    {"line_number": line_num, "diff": diff} for line_num, diff in differences
                ],
            }

            with open(args.output, "w") as f:
                json.dump(output_data, f, indent=2)

            print(f"\nDetailed differences saved to: {args.output}")

        print("=" * 80)

        return 0 if not differences else 1

    except KeyboardInterrupt:
        print("\n\nInterrupted by user")
        return 1
    except Exception as e:
        print(f"\nError: {e}")
        import traceback

        traceback.print_exc()
        return 1


if __name__ == "__main__":
    exit(main())
