import sys
import json
from datetime import datetime
from argparse import ArgumentParser

from orionclient.session import APISession
from orionclient.exceptions import OrionError
from orionclient.types import Dataset, WorkFloeSpec
from orionclient.helpers.parameterize import parameterize_workfloe

WORKFLOE_TITLE = "Classic OMEGA"


def run_benchmarks(args):
    try:
        dataset = APISession.get_resource(Dataset, args.dataset_id)
    except OrionError as e:
        print(e)
        sys.exit(1)
    filters = {"title": WORKFLOE_TITLE}
    workfloes = list(APISession.list_resources(WorkFloeSpec, filters=filters))
    if not len(workfloes):
        print(f"Unable to find a workfloe named {WORKFLOE_TITLE}")
        sys.exit(1)
    workfloe = workfloes[0]
    # Refresh to get the specification
    APISession.refresh_resource(workfloe)
    print(f"Benchmarking Workfloe {workfloe.title}, id:{workfloe.id:d}")
    default_params = {"promoted": {"in": dataset.id, "out": args.output_dataset}}
    # Get the cubes in the specification and get default values for parameter
    # being benchmarked
    cubes = workfloe.specification["cubes"]
    default_values = {}
    for cube in cubes:
        for param in cube["parameters"]:
            if param["name"] == args.benchmark:
                default_values[cube["name"]] = {
                    "default": param["default"],
                    "max_value": param["max_value"],
                }
    parameters = []
    for x in range(args.steps):
        cube_params = {}
        for key, item in default_values.items():
            cube_params[key] = {args.benchmark: item["default"] * (x + 1)}
            if item["max_value"] is not None:
                cube_params[key][args.benchmark] = min(
                    cube_params[key][args.benchmark], item["max_value"]
                )
        parameters.append({"cube": cube_params})
    jobs = parameterize_workfloe(
        workfloe,
        f"{WORKFLOE_TITLE} Benchmark",
        default_params,
        parameters,
        parallel=args.parallel,
    )
    for job in jobs:
        header = f"Job {job.id}"
        print(header)
        print("-" * len(header))
        print(json.dumps(job.parameters, indent=2))
        started = datetime.strptime(job.started, "%Y-%m-%d %H:%M")
        ended = datetime.strptime(job.finished, "%Y-%m-%d %H:%M")
        print(f"Took {ended - started}")
        if not job.success:
            print(f"Failure - {job.reason}")
        else:
            print("Success")


def main():
    parser = ArgumentParser(description=f"Benchmark {WORKFLOE_TITLE}")
    parser.add_argument("dataset_id")
    parser.add_argument("output_dataset")
    parser.add_argument(
        "--benchmark",
        default="buffer_size",
        choices=["buffer_size", "item_count"],
        help="Cube parameter to benchmark",
    )
    parser.add_argument(
        "--steps",
        default=5,
        type=int,
        help="Number of steps to run, each step is default times the step number",
    )
    parser.add_argument(
        "--parallel", action="store_true", help="Run benchmarking floes in parallel"
    )
    args = parser.parse_args()
    run_benchmarks(args)


if __name__ == "__main__":
    main()
