#!/usr/bin/env python3

import argparse
import collections
import concurrent.futures
import csv
import functools
import json
import os
import subprocess

# ========== Args ==========


def parse_args():
    parser = argparse.ArgumentParser(description="solve all instances")
    parser.add_argument(
        "-i",
        "--input_dir",
        help="Path where all instances are located",
        default="./1995_bischoff_and_ratcliff_and_1999_davies_and_bischoff",
    )
    parser.add_argument(
        "-o",
        "--output_dir",
        help="Path of the directory where all results will be saved",
        default="./solutions_output",
    )
    parser.add_argument(
        "-t",
        "--maximum_running_time_seconds",
        help="maximum running time for the allocation, in seconds",
        default="30",
    )
    parser.add_argument(
        "-c",
        "--configurations",
        help="csv file containing all configurations to test",
        default=None,
    )
    parser.add_argument(
        "--commit",
        help="the short hash of the commit/image to use",
        default="latest",
    )
    parser.add_argument(
        "--reduced",
        help="Enabled: solve only the instances 1 to 10 (included) of all sets; Disabled: solve all instances available",
        action=argparse.BooleanOptionalAction,
    )
    parser.add_argument(
        "--local",
        help="use locally compiled app instead of docker image",
        action=argparse.BooleanOptionalAction,
    )
    parser.add_argument(
        "--app",
        help="Path to compiled app",
        default="../slopp-heuristic/Release/app/allocate",
    )
    parser.add_argument(
        "--dry-run",
        help="show the commands it will execute and exit",
        action=argparse.BooleanOptionalAction,
    )
    parser.add_argument(
        "-j",
        "--threads",
        help="number of threads to use / number of jobs to run in parallel",
        type=int,
        default=os.cpu_count(),
    )
    args = parser.parse_args()
    return args


cli_args = parse_args()


# ========== Allocate ==========

# ---------- Configuration definitions ----------


class Configuration:
    properties = (
        # (name, type, default_value)
        ("maximum_running_time_seconds", int, 30),
        ("id", int, 0),
        ("alpha", float, 4.0),
        ("beta", float, 1.0),
        ("gamma", float, 0.2),
        ("delta", float, 0.04),
        ("skip_type", str, "strict"),
        ("initial_branching_factor", int, 4),
        ("maximum_branching_factor", int, 10000),
        ("block_minimum_volume_ratio", float, -1.0),
        ("maximum_number_of_blocks", int, 10000),
        ("alpha_2", float, 4.0),
        ("beta_2", float, 1.0),
        ("gamma_2", float, 0.2),
        ("delta_2", float, 0.04),
        ("configuration_switch_threshold", float, 1.01),
    )

    @staticmethod
    def property_names():
        return (p[0] for p in Configuration.properties)

    def __init__(self):
        self.values = {}
        # for name, type, default_value in Configuration.properties:
        #     self.values[name] = type(default_value) if default_value else None
        return

    @staticmethod
    def from_file(file_path: str):
        if file_path is None:
            config = Configuration()
            config.values["id"] = 0
            return [config]
        else:
            configurations = []
            id = 0
            with open(file_path, "r") as file:
                data = csv.DictReader(file)
                for row in data:
                    config = Configuration()
                    config.values["id"] = id
                    for name, type, default_value in Configuration.properties:
                        if name not in row:
                            continue
                        if row[name] is None:
                            continue
                        if row[name] == "":
                            continue
                        config.values[name] = type(row[name])
                    configurations.append(config)
                    id += 1
            return configurations

    def get(self, name: str):
        return self.values[name]

    def to_cli_args(self) -> list[str]:
        cli_args = []
        if "maximum_running_time_seconds" not in self.values.keys():
            default_maximum_running_time_seconds = 30
            cli_args.append(f"--maximum_running_time_seconds={default_maximum_running_time_seconds}")
        for name, value in self.values.items():
            if name == "id":
                continue
            if value is None:
                continue
            cli_args.append(f"--{name}={value}")
        return cli_args

    def to_csv(self) -> list[str]:
        csv = []
        for name in Configuration.property_names():
            if name not in self.values:
                csv.append("")
            else:
                csv.append(str(self.values[name]))
        return csv


# ---------- Instance definitions ----------


class Instance:
    def __init__(self, instances_dir, instance_set, instance_number):
        self.instance_set = instance_set
        self.instance_number = instance_number
        self.set_dir = f"set_{instance_set}"
        self.name = f"{instance_number}.json"
        self.relative_path = os.path.join(self.set_dir, self.name)
        self.path = os.path.join(instances_dir, self.relative_path)
        self.valid = os.path.isfile(self.path)

    def __repr__(self):
        return str(self.__dict__)


def find_instances(instances_dir, instance_sets, instance_numbers):
    instances = [
        Instance(instances_dir, instance_set, instance_number)
        for instance_set in instance_sets
        for instance_number in instance_numbers
    ]
    # check that they are all valid
    for instance in instances:
        if not instance.valid:
            raise RuntimeError(f"instance does not exist:\n{instance}")
    return instances


# ---------- Script variables ----------

# directories
input_dir = os.path.realpath(cli_args.input_dir)
output_dir = os.path.realpath(cli_args.output_dir)
os.makedirs(output_dir, exist_ok=True)
container_input_dir = "/tmp/input"
container_output_dir = "/tmp/output"

# instances and configurations
instance_sets = range(0, 16)
instance_numbers = range(1, 11) if cli_args.reduced else range(1, 101)
instances = find_instances(input_dir, instance_sets, instance_numbers)
configurations = Configuration.from_file(cli_args.configurations)
for instance in instances:
    for configuration in configurations:
        # make all directories where all output data will be saved
        instance_output_dir = os.path.join(
            output_dir, str(configuration.get("id")), instance.set_dir
        )
        if not cli_args.dry_run:
            os.makedirs(instance_output_dir, exist_ok=True)

# logs
log_file = os.path.join(output_dir, "solve.log")

# container image
registry = "lucasguesserts/private"
image = f"{registry}:slopp-app-{cli_args.commit}"

# system resources
memory_fraction_usage = 0.8
total_available_memory_MB = (
    os.sysconf("SC_PAGE_SIZE") * os.sysconf("SC_PHYS_PAGES") / (1024.0**2)
)
memory_usage_limit_per_container_MB = int(
    memory_fraction_usage * total_available_memory_MB / cli_args.threads
)
user_id = os.getuid()
assert 1 <= cli_args.threads and cli_args.threads <= os.cpu_count(), f"the number of threads/jobs must be between {1} and {os.cpu_count()}, but it is {cli_args.threads}"
assert 6 <= memory_usage_limit_per_container_MB
assert memory_usage_limit_per_container_MB <= total_available_memory_MB


# ---------- Solver definitions ----------


def solve_instance(arg: tuple[Instance, Configuration]):
    instance, configuration = arg
    if cli_args.local:
        command = [
            cli_args.app,
            *configuration.to_cli_args(),
            f"--output-file={output_dir}/{configuration.get('id')}/{instance.relative_path}",
            f"{input_dir}/{instance.relative_path}",
        ]
    else:
        command = [
            "docker",
            "container",
            "run",
            "--rm",
            "-t",
            f"--user={user_id}:{user_id}",
            "-v",
            f"{input_dir}:{container_input_dir}:z",
            "-v",
            f"{output_dir}:{container_output_dir}:z",
            f"--memory={memory_usage_limit_per_container_MB}m",
            f"--memory-swap={memory_usage_limit_per_container_MB}m",
            f"{image}",
            *configuration.to_cli_args(),
            f"--output-file={container_output_dir}/{configuration.get('id')}/{instance.relative_path}",
            f"{container_input_dir}/{instance.relative_path}",
        ]
    if cli_args.dry_run:
        command = " ".join(command)
        print(command, "\n")
        ret = {
            "command": command,
            "return_code": "dry-run",
            "stdout": "dry-run",
            "stderr": "dry-run",
        }
    else:
        print(
            f"--- start - set {instance.instance_set} - instance {instance.instance_number} - configuration {configuration.get('id')}"
        )
        ret = subprocess.run(
            command,
            capture_output=True,
            text=True,
        )
        ret = {
            "command": " ".join(ret.args),
            "return_code": ret.returncode,
            "stdout": ret.stdout,
            "stderr": ret.stderr,
        }
        print(
            f"=== finish - set {instance.instance_set} - instance {instance.instance_number}"
        )
    return ret


# Function to execute tasks in batches of N parallel workers
def run_parallel(task_function, task_args, number_of_parallel_tasks):
    if cli_args.dry_run:
        return_values = [task_function(args) for args in task_args]
        return return_values
    return_values = []
    with concurrent.futures.ProcessPoolExecutor(
        max_workers=number_of_parallel_tasks
    ) as executor:
        futures = {executor.submit(task_function, arg): arg for arg in task_args}
        for future in concurrent.futures.as_completed(futures):
            result = future.result()
            return_values.append(result)
    return return_values


def check_all_outputs_exist(instances, configurations):
    if cli_args.dry_run:
        return
    for instance in instances:
        for configuration in configurations:
            output_file = (
                f"{output_dir}/{configuration.get('id')}/{instance.relative_path}"
            )
            if not os.path.isfile(output_file):
                raise RuntimeError(
                    f"file '{output_file} has not been written. Check logs for more info"
                )
    return


def allocate():
    args = [
        (instance, configuration)
        for configuration in configurations
        for instance in instances
    ]
    logs = run_parallel(solve_instance, args, cli_args.threads)
    with open(log_file, "w") as file:
        json.dump(logs, file, indent=2, ensure_ascii=True)
    check_all_outputs_exist(instances, configurations)
    return


# ========== join ==========


def join_output_files(directory, output_file):
    if cli_args.dry_run:
        return
    output = {}
    instances_set_directories = os.listdir(directory)
    instances_set_directories = [
        os.path.join(directory, dir) for dir in instances_set_directories
    ]
    instances_set_directories = list(
        filter(lambda e: os.path.isdir(e), instances_set_directories)
    )
    instances_set_directories.sort(key=lambda d: int(d.split("_")[-1]))
    for dir in instances_set_directories:
        set_number = dir.split("_")[-1]
        set_output = {}
        file_name_list = os.listdir(dir)
        file_name_list.sort(key=lambda f: int(f.split(".")[0]))
        for file_name in file_name_list:
            instance_number = file_name.split(".")[0]
            with open(os.path.join(dir, file_name), "r") as file:
                set_output[instance_number] = json.load(file)
        output[set_number] = set_output
    with open(output_file, "w") as file:
        json.dump(output, file, indent=2, ensure_ascii=True)
    return


# ========== analyse ==========


def make_history(file_path):
    history = []
    with open(file_path, "r") as file:
        data = json.load(file)
    for instance_set in data:
        for instance_number in data[instance_set]:
            instance_data = data[instance_set][instance_number]["appendix"]
            # check that "solution_history" exists and it has something
            # if not, use "packing_statistics"
            solution_history_present = "solution_history" in instance_data
            solution_history_not_empty = (
                len(instance_data["solution_history"]) > 0
                if solution_history_present
                else False
            )
            if solution_history_not_empty:
                solution_history = instance_data["solution_history"]
                for entry in solution_history:
                    history.append(
                        [
                            int(instance_set),
                            int(instance_number),
                            float(entry["time"]),
                            float(100 * entry["volume_usage"]) / (587 * 233 * 220),
                        ]
                    )
            else:
                history.append(
                    [
                        int(instance_set),
                        int(instance_number),
                        float(0.01),
                        float(
                            instance_data["packing_statistics"][
                                "volume_fraction_relative_to_large_object_volume"
                            ]
                        ),
                    ]
                )
    history.sort()
    return history


def collect_volume_usage(history, time_limit):
    # group by key = (instance set, instance number)
    # for each group, keep only the max volume usage
    # set the volume usage of keys not set to zero
    volume_usages = {}
    key_set = set()  # store all keys
    for entry in history:
        key = (entry[0], entry[1])
        key_set.add(key)
        time = entry[2]
        if time > time_limit:
            continue
        volume_usage = entry[3]
        if key not in volume_usages:
            volume_usages[key] = volume_usage
        else:
            volume_usages[key] = max(volume_usages[key], volume_usage)
    # set values for all keys not added
    for key in key_set:
        if key not in volume_usages:
            volume_usages[key] = 0.0
    # transform into list
    volume_usages = [(*key, value) for key, value in volume_usages.items()]
    # sort
    volume_usages.sort()
    return volume_usages


def write_volume_usage_csv(volume_usages, file_path):
    with open(file_path, "w") as csvfile:
        writer = csv.writer(
            csvfile, delimiter=",", lineterminator="\n", quoting=csv.QUOTE_NONE
        )
        writer.writerow(["instance_set", "instance_number", "volume_usage"])  # header
        for entry in volume_usages:
            writer.writerow(entry)  # entries
    return


def compute_summary(volume_usages):
    groups = {}
    # group entries of the same instance set
    for instance_set, instance_number, volume_usage in volume_usages:
        key = str(instance_set)
        if key in groups:
            groups[key].append(volume_usage)
        else:
            groups[key] = [volume_usage]
    # group entries that satisfy specific conditions
    ## 1-7
    key = "1-7"
    groups[key] = []
    for instance_set, instance_number, volume_usage in volume_usages:
        if 1 <= instance_set and instance_set <= 7:
            groups[key].append(volume_usage)
    ## 8-15
    key = "8-15"
    groups[key] = []
    for instance_set, instance_number, volume_usage in volume_usages:
        if 8 <= instance_set and instance_set <= 15:
            groups[key].append(volume_usage)
    ## 1-15
    key = "1-15"
    groups[key] = []
    for instance_set, instance_number, volume_usage in volume_usages:
        if 1 <= instance_set and instance_set <= 15:
            groups[key].append(volume_usage)
    # compute averages
    summary = {}
    for key, values in groups.items():
        if len(values) == 0:
            print(values)
        summary[key] = round(sum(values) / len(values), ndigits=2)
    return summary


def write_summary_csv(summary, file_path):
    with open(file_path, "w") as csvfile:
        writer = csv.writer(
            csvfile, delimiter=",", lineterminator="\n", quoting=csv.QUOTE_NONE
        )
        writer.writerow(["instance_set", "volume_usage"])  # header
        for entry in summary.items():
            writer.writerow(entry)  # entries
    return


def analyse(
    input_file, time_limit, individual_results_file_name, summary_results_file_name
):
    if cli_args.dry_run:
        return
    history = make_history(input_file)
    volume_usages = collect_volume_usage(history, float(time_limit))
    summary = compute_summary(volume_usages)
    write_volume_usage_csv(volume_usages, individual_results_file_name)
    write_summary_csv(summary, summary_results_file_name)
    return


def volume_usage_by_configuration():
    if cli_args.dry_run:
        return None
    # initialize main table: add all instance sets and numbers
    main_table = []
    for instance in instances:
        main_table.append([instance.instance_set, instance.instance_number])
    main_table.sort()
    # load the "extracted_data.csv" for each configuration
    # and append the volume usage to the main table
    for configuration in configurations:
        configuration_table = []
        configuration_results_file = os.path.join(
            output_dir, str(configuration.get("id")), "extracted_data.csv"
        )
        with open(configuration_results_file, "r") as file:
            data = csv.DictReader(file)
            for row in data:
                configuration_table.append(
                    [
                        # use instance_set and instance_number for sorting
                        # to guarantee that the entries of all tables match
                        int(row["instance_set"]),
                        int(row["instance_number"]),
                        float(row["volume_usage"]),
                    ]
                )
        configuration_table.sort()
        configuration_table = list(
            map(lambda e: e[2], configuration_table)
        )  # remove first two entries
        for main_entry, configuration_entry in zip(main_table, configuration_table):
            main_entry.append(configuration_entry)
    # write csv file
    file_path = os.path.join(output_dir, "extracted_data_by_configuration.csv")
    with open(file_path, "w") as file:
        writer = csv.writer(
            file, delimiter=",", lineterminator="\n", quoting=csv.QUOTE_NONE
        )
        header = ["instance_set", "instance_number"] + [
            f"{c.get('id')}" for c in configurations
        ]
        writer.writerow(header)
        for entry in main_table:
            writer.writerow(entry)
    return main_table


def volume_usage_by_configuration_summary(table: list[list]):
    if cli_args.dry_run:
        return
    groups = {}
    # group entries of the same instance set
    for instance_set, instance_number, *volume_usage in table:
        key = str(instance_set)
        if key in groups:
            groups[key].append(volume_usage)
        else:
            groups[key] = [volume_usage]
    # group entries that satisfy specific conditions
    ## 1-7
    key = "1-7"
    groups[key] = []
    for instance_set, instance_number, *volume_usage in table:
        if 1 <= instance_set and instance_set <= 7:
            groups[key].append(volume_usage)
    ## 8-15
    key = "8-15"
    groups[key] = []
    for instance_set, instance_number, *volume_usage in table:
        if 8 <= instance_set and instance_set <= 15:
            groups[key].append(volume_usage)
    ## 1-15
    key = "1-15"
    groups[key] = []
    for instance_set, instance_number, *volume_usage in table:
        if 1 <= instance_set and instance_set <= 15:
            groups[key].append(volume_usage)
    # compute averages
    summary = {}
    for key, values in groups.items():
        sums = list(
            functools.reduce(lambda acc, vs: [a + v for a, v in zip(acc, vs)], values)
        )
        sums = [s / len(values) for s in sums]
        sums = [round(s, ndigits=2) for s in sums]
        summary[key] = sums
    # write
    file_path = os.path.join(output_dir, "extracted_data_summary_by_configuration.csv")
    with open(file_path, "w") as csvfile:
        writer = csv.writer(
            csvfile, delimiter=",", lineterminator="\n", quoting=csv.QUOTE_NONE
        )
        header = ["instance_set"] + [f"{c.get('id')}" for c in configurations]
        writer.writerow(header)
        for key, values in summary.items():
            writer.writerow([key] + values)  # entries
    return


def write_configurations():
    file_path = os.path.join(output_dir, "configurations.csv")
    with open(file_path, "w") as file:
        writer = csv.writer(
            file, delimiter=",", lineterminator="\n", quoting=csv.QUOTE_NONE
        )
        header = Configuration.property_names()
        writer.writerow(header)
        for configuration in configurations:
            writer.writerow(configuration.to_csv())
    return


if __name__ == "__main__":
    allocate()
    for configuration in configurations:
        local_output_dir = os.path.join(output_dir, f"{configuration.get('id')}")
        all_file_path = os.path.join(local_output_dir, "all.json")
        join_output_files(local_output_dir, all_file_path)
        analyse(
            all_file_path,
            cli_args.maximum_running_time_seconds,
            os.path.join(local_output_dir, "extracted_data.csv"),
            os.path.join(local_output_dir, "extracted_data_summary.csv"),
        )
    write_configurations()
    volume_usage_by_configuration_table = volume_usage_by_configuration()
    volume_usage_by_configuration_summary(volume_usage_by_configuration_table)
