import os
import re
import argparse
from collections import defaultdict
from obspy import read_inventory, UTCDateTime

# Standard gravity for conversion from m/s^2 to g
G_TO_MS2 = 9.80665
DIFF_THRESHOLD = 1.0  # Percentage threshold for recording a discrepancy

def parse_cosmos_v0c(v0c_filepath):
    """Parses a COSMOS v0c file silently to extract Real Headers, SCNL, and Origin time."""
    rh22, rh42 = None, None
    net, sta, loc, cha = "*", "*", "*", "*"
    origin_time = None
    real_headers = []
    
    try:
        with open(v0c_filepath, 'r', encoding='utf-8', errors='ignore') as f:
            lines = f.readlines()
    except FileNotFoundError:
        return None, None, net, sta, loc, cha, origin_time

    in_real_section = False
    for line in lines:
        line_upper = line.upper()
        
        if 'ORIGIN:' in line_upper and origin_time is None:
            match = re.search(r'Origin:\s*(\d{4}/\d{2}/\d{2}\s+\d{2}:\d{2}:\d{2})', line)
            if match:
                try:
                    origin_time = UTCDateTime.strptime(match.group(1), "%Y/%m/%d %H:%M:%S")
                except Exception:
                    pass

        if 'REAL VAL' in line_upper or 'REAL-HEADER' in line_upper or 'REAL HEADER' in line_upper:
            in_real_section = True
            continue

        if in_real_section:
            if 'COMMENT' in line_upper or 'TEXT' in line_upper or 'END' in line_upper or 'DATA' in line_upper:
                in_real_section = False
            else:
                floats = re.findall(r'[-+]?(?:\d*\.\d+|\d+\.)(?:[eE][-+]?\d+)?', line)
                real_headers.extend([float(x) for x in floats])

        if '<SCNL>' in line_upper:
            match = re.search(r'<SCNL>([^\s]+)', line)
            if match:
                parts = match.group(1).split('.')
                if len(parts) >= 4:
                    sta, cha, net, loc = parts[0], parts[1], parts[2], parts[3]
                    if loc == '--': loc = ""

    if len(real_headers) >= 42:
        rh22 = real_headers[21]
        rh42 = real_headers[41]

    return rh22, rh42, net, sta, loc, cha, origin_time

def get_stationxml_stages(xml_filepath, network, station, location, channel, origin_time):
    """Parses a StationXML file to extract Stage 1 (Sensor) and Digitizer stage (V -> COUNTS)."""
    try:
        inv = read_inventory(xml_filepath)
        inv_selected = inv.select(network=network, station=station, location=location, channel=channel, time=origin_time)
        
        if not inv_selected:
            return None, None, None, None

        response = inv_selected[0][0][0].response
        stage1_gain, stage1_units, digitizer_gain, digitizer_stage_num = None, None, None, None

        for stage in response.response_stages:
            if stage.stage_sequence_number == 1:
                stage1_gain = stage.stage_gain
                stage1_units = stage.input_units
            
            if stage.input_units and stage.output_units:
                if stage.input_units.upper() == "V" and stage.output_units.upper() == "COUNTS":
                    digitizer_gain = stage.stage_gain
                    digitizer_stage_num = stage.stage_sequence_number

        return stage1_gain, stage1_units, digitizer_gain, digitizer_stage_num
    except Exception:
        return None, None, None, None

def check_file_for_discrepancy(v0c_file, xml_base_dir, show_details):
    """
    Checks for discrepancies. 
    Returns (net_sta, flags_set, details_string).
    Returns (None, set(), "") if it's a perfect match or invalid file.
    """
    rh22, rh42, net, sta, loc, cha, origin_time = parse_cosmos_v0c(v0c_file)
    
    if None in (rh22, rh42) or net == "*" or sta == "*":
        return None, set(), ""
        
    net_sta = f"{net}.{sta}"
    filename = os.path.basename(v0c_file)
    details = []
    flags = set()

    xml_file = os.path.join(xml_base_dir, f"{net}.info", f"{net}.FDSN.xml", f"{net}.{sta}.xml")

    # Missing XML
    if not os.path.exists(xml_file):
        flags.add("XML_MISSING")
        if show_details:
            return net_sta, flags, f"    [File: {filename} | Cha: {cha}] Error: StationXML not found -> {xml_file}"
        return net_sta, flags, ""

    s1_gain, s1_units, digi_gain, digi_stage_num = get_stationxml_stages(xml_file, net, sta, loc, cha, origin_time)

    # Missing Stages in XML
    if s1_gain is None or digi_gain is None:
        flags.add("METADATA_MISSING")
        if show_details:
            return net_sta, flags, f"    [File: {filename} | Cha: {cha}] Error: Missing Stage 1 or Digitizer metadata in XML."
        return net_sta, flags, ""

    # Conversions
    s1_gain_converted = s1_gain
    conversion_note_s1 = ""
    if s1_units and "M/S**2" in s1_units.upper():
        s1_gain_converted = s1_gain * G_TO_MS2
        conversion_note_s1 = ""
    
    digi_gain_converted = (1.0 / digi_gain) * 1_000_000 if digi_gain != 0 else 0

    # Differences
    s1_diff = abs(s1_gain_converted - rh42) / max(abs(s1_gain_converted), 1e-12) * 100
    s2_diff = abs(digi_gain_converted - rh22) / max(abs(digi_gain_converted), 1e-12) * 100

    if s1_diff > DIFF_THRESHOLD:
        flags.add("RH42")
    if s2_diff > DIFF_THRESHOLD:
        flags.add("RH22")

    if flags:
        if show_details:
            details.append(f"    [File: {filename} | Cha: {cha} | Epoch: {origin_time}]")
            
            if "RH42" in flags:
                details.append(f"      --- SENSOR SENSITIVITY (RH42) ---")
                details.append(f"      XML Stage 1     : {s1_gain:.6e} V / {s1_units}")
                details.append(f"      XML (Converted) : {s1_gain_converted:.6e} V/g {conversion_note_s1}")
                details.append(f"      COSMOS RH42     : {rh42:.6e} V/g")
                details.append(f"      Difference      : {s1_diff:.4f} %")
                
            if "RH22" in flags:
                details.append(f"      --- RECORDER LSB (RH22) ---")
                details.append(f"      XML Stage {digi_stage_num: <2}    : {digi_gain:.6e} Counts/V")
                details.append(f"      XML (Converted) : {digi_gain_converted:.6e} uV/Count")
                details.append(f"      COSMOS RH22     : {rh22:.6e} uV/Count")
                details.append(f"      Difference      : {s2_diff:.4f} %")
                
            return net_sta, flags, "\n".join(details)
        return net_sta, flags, ""

    return None, set(), ""

def main():
    parser = argparse.ArgumentParser(description="Batch compare COSMOS Real Headers with StationXML Stage Gains.")
    parser.add_argument("path", help="Directory containing event folders (e.g., /work/ftp/outgoing/V0DataFiles/)")
    parser.add_argument("--xml-base-dir", default="/work/pub/doc", help="Base directory for StationXML files")
    parser.add_argument("--details", action="store_true", help="Print detailed breakdown for discrepancies")
    
    args = parser.parse_args()

    if not os.path.exists(args.path):
        print(f"Error: Path '{args.path}' does not exist.")
        return

    print(f"Scanning directory: {args.path}")

    # Structures to hold flags and details separately
    discrepancy_flags = defaultdict(lambda: defaultdict(set))
    discrepancy_details = defaultdict(lambda: defaultdict(list))

    # Walk through the directory tree
    for root, dirs, files in os.walk(args.path):
        v0c_files = [f for f in files if f.endswith('.v0c')]
        
        if v0c_files:
            event_name = os.path.basename(root)
            
            for v0c in v0c_files:
                filepath = os.path.join(root, v0c)
                net_sta, flags, details = check_file_for_discrepancy(filepath, args.xml_base_dir, args.details)
                
                if net_sta and flags:
                    discrepancy_flags[event_name][net_sta].update(flags)
                    if details:
                        discrepancy_details[event_name][net_sta].append(details)

    # --- Print Summary Report ---
    if not discrepancy_flags:
        print("All processed files matched StationXML metadata within the 1% threshold.")
    else:
        print("DISCREPANCIES FOUND (> 1% diff or missing XML metadata)")
        print("="*65)
        
        for event in sorted(discrepancy_flags.keys()):
            print(f"Event: {event}")
            
            for station in sorted(discrepancy_flags[event].keys()):
                # Format the tags (e.g., [RH22] [RH42])
                flags_str = " ".join([f"[{f}]" for f in sorted(discrepancy_flags[event][station])])
                print(f"  - {station} {flags_str}")
                
                # If --details was passed, print out the gathered string payloads
                if args.details:
                    for detail_str in discrepancy_details[event][station]:
                        if detail_str:
                            print(detail_str)
                    print("") # Extra spacing between stations if details are shown
            print("") # Extra spacing between events

if __name__ == "__main__":
    main()
