#!/usr/bin/env python3

import os
import re
import csv
import argparse
import sys

def extract_parameters_from_log(log_file):
    try:
        parameters = []
        with open(log_file, "r") as file:
            lines = file.readlines()

        for i, line in enumerate(lines):
            if "# Best configurations as commandlines" in line:
                for config_line in lines[i + 1 :]:
                    if not config_line.strip():
                        break
                    match = re.match(
                        r"(\d+)\s+--alpha=([0-9.]+)\s+--beta=([0-9.]+)\s+--gamma=([0-9.]+)\s+--delta=([0-9.]+)",
                        config_line,
                    )
                    if match:
                        parameters.append(
                            {
                                "config_id": int(match.group(1)),
                                "alpha": float(match.group(2)),
                                "beta":  float(match.group(3)),
                                "gamma": float(match.group(4)),
                                "delta": float(match.group(5)),
                            }
                        )
        if not parameters:
            raise ValueError(f"No valid configurations found in {log_file}.")

        return parameters
    except FileNotFoundError:
        raise FileNotFoundError(f"Log file {log_file} not found.")
    except Exception as e:
        raise Exception(f"Error processing log file {log_file}: {str(e)}")

def process_file(input_file, output_csv):
    try:
        parameters = extract_parameters_from_log(input_file)

        with open(output_csv, "w", newline="") as csvfile:
            fieldnames = ["alpha", "beta", "gamma", "delta"]
            writer = csv.DictWriter(csvfile, fieldnames=fieldnames)
            writer.writeheader()

            for param in parameters:
                writer.writerow(
                    {
                        "alpha": "{:.4f}".format(param["alpha"]),
                        "beta": "{:.4f}".format(param["beta"]),
                        "gamma": "{:.4f}".format(param["gamma"]),
                        "delta": "{:.4f}".format(param["delta"]),
                    }
                )
        print(f"Processing completed successfully. Data written to {output_csv}")
    except Exception as e:
        print(f"Error processing file {input_file}: {str(e)}", file=sys.stderr)

def process_directory(input_dir, output_csv, include_instance_info):
    try:
        if not os.path.isdir(input_dir):
            raise NotADirectoryError(f"Input directory {input_dir} does not exist or is not a directory.")

        with open(output_csv, "w", newline="") as csvfile:
            fieldnames = ["alpha", "beta", "gamma", "delta"]
            if include_instance_info:
                fieldnames = ["instance_set", "instance_number"] + fieldnames

            writer = csv.DictWriter(csvfile, fieldnames=fieldnames)
            writer.writeheader()

            for subdir in os.listdir(input_dir):
                subdir_path = os.path.join(input_dir, subdir)

                if os.path.isdir(subdir_path):
                    match = re.match(r"set_(\d+)_instance_(\d+)", subdir)
                    if match:
                        instance_set = match.group(1)
                        instance_number = match.group(2)

                        log_file = os.path.join(subdir_path, "irace.log")
                        if os.path.isfile(log_file):
                            try:
                                parameters = extract_parameters_from_log(log_file)
                                for param in parameters:
                                    row = {
                                        "alpha": "{:.4f}".format(param["alpha"]),
                                        "beta": "{:.4f}".format(param["beta"]),
                                        "gamma": "{:.4f}".format(param["gamma"]),
                                        "delta": "{:.4f}".format(param["delta"]),
                                    }
                                    if include_instance_info:
                                        row["instance_set"] = instance_set
                                        row["instance_number"] = instance_number

                                    writer.writerow(row)
                            except Exception as e:
                                print(f"Warning: Skipping {log_file}. Reason: {str(e)}", file=sys.stderr)
                        else:
                            print(f"Warning: Log file {log_file} not found in {subdir_path}. Skipping this directory.", file=sys.stderr)
    except Exception as e:
        raise Exception(f"Error processing directory {input_dir}: {str(e)}")

def main():
    parser = argparse.ArgumentParser(description="Extract optimal parameters from irace log files.")
    parser.add_argument(
        "--input",
        type=str,
        required=True,
        help="Path to a single irace.log file or a directory containing subdirectories with irace.log files",
    )
    parser.add_argument(
        "--output-file",
        type=str,
        required=True,
        help="CSV file to store extracted parameters",
    )
    parser.add_argument(
        "--include-instance-info",
        action="store_true",
        help="Include instance set and instance number columns in the CSV output (only for directory processing)",
    )

    args = parser.parse_args()

    try:
        if os.path.isdir(args.input):
            process_directory(args.input, args.output_file, args.include_instance_info)
        elif os.path.isfile(args.input):
            process_file(args.input, args.output_file)
        else:
            raise ValueError(f"Invalid input path: {args.input}. It must be an existing file or directory.")
    except Exception as e:
        print(f"Error: {str(e)}", file=sys.stderr)

if __name__ == "__main__":
    main()
