#!/usr/bin/env python3
import argparse
import csv
import gzip
import html
import io
import json
import os
import re
import sys
import time
import urllib.error
import urllib.request

GIB = 1024 ** 3

# Cloud key and display name; the fetcher is the matching get_<key> function.
CLOUDS = (
    ("aws", "Amazon Web Services"),
    ("azure", "Microsoft Azure"),
    ("digitalocean", "DigitalOcean"),
    ("gce", "Google Cloud"),
    ("hetzner", "Hetzner Cloud"),
    ("linode", "Linode (Akamai)"),
    ("ovh", "OVH"),
    ("scaleway", "Scaleway"),
    ("vultr", "Vultr"),
)

# Refuse to overwrite a file when a source suddenly returns far fewer entries.
MIN_COUNT = {
    "aws": 400,
    "azure": 400,
    "gce": 100,
    "digitalocean": 40,
    "hetzner": 15,
    "linode": 30,
    "ovh": 40,
    "scaleway": 30,
    "vultr": 40,
}

# AWS families built around accelerators or Apple hardware.
AWS_ACCEL_FAMILY = re.compile(r"^(p\d|g\d|gr\d|inf\d|trn\d|dl\d|vt\d|f\d|mac\d)")

# AWS storage descriptions, e.g. "1 x 950 NVMe SSD" or "2 x 1900 SSD".
AWS_STORAGE = re.compile(r"^(\d+) x ([\d.]+)")


def entry(cpu, mem, arch="x86_64", disk=None):
    data = {"arch": arch, "cpu": float(cpu), "mem": float(mem)}
    if disk:
        data["disk"] = int(disk)

    return data


def fetch(url, headers=None, retries=3):
    """Fetch a URL and return a binary file object, retrying on failure."""
    req = urllib.request.Request(url, headers=headers or {})

    for attempt in range(retries):
        try:
            return urllib.request.urlopen(req, timeout=300)
        except (urllib.error.URLError, TimeoutError):
            if attempt == retries - 1:
                raise

            time.sleep(5 * (attempt + 1))


def fetch_json(url, headers=None):
    hdrs = {"Accept-Encoding": "gzip", **(headers or {})}
    resp = fetch(url, headers=hdrs)
    data = resp.read()
    if resp.headers.get("Content-Encoding") == "gzip" or data[:2] == b"\x1f\x8b":
        data = gzip.decompress(data)

    return json.loads(data)


def get_aws():
    # Stream the us-east-1 EC2 offer file (~300MB); it covers nearly all types.
    url = "https://pricing.us-east-1.amazonaws.com/offers/v1.0/aws/AmazonEC2/current/us-east-1/index.csv"
    resp = fetch(url)
    reader = csv.reader(io.TextIOWrapper(resp, encoding="utf-8", newline=""))

    # Skip the metadata preamble before the real header row.
    header = None
    for row in reader:
        if "Instance Type" in row:
            header = row
            break

    if header is None:
        raise RuntimeError("aws: header row not found")

    col_family = header.index("Product Family")
    col_type = header.index("Instance Type")
    col_cpu = header.index("vCPU")
    col_mem = header.index("Memory")
    col_gpu = header.index("GPU") if "GPU" in header else None
    col_proc = header.index("Physical Processor") if "Physical Processor" in header else None
    col_disk = header.index("Storage") if "Storage" in header else None

    types = {}
    for row in reader:
        if len(row) <= max(col_family, col_type, col_cpu, col_mem):
            continue

        if row[col_family] not in ("Compute Instance", "Compute Instance (bare metal)"):
            continue

        name = row[col_type]
        if not name or name in types:
            continue

        if AWS_ACCEL_FAMILY.match(name.split(".")[0]):
            continue

        if col_gpu is not None and row[col_gpu] not in ("", "0", "NA"):
            continue

        mem = row[col_mem].replace(",", "").strip()
        if not mem.endswith(" GiB"):
            continue

        arch = "x86_64"
        if col_proc is not None and "Graviton" in row[col_proc]:
            arch = "aarch64"

        disk = None
        if col_disk is not None:
            match = AWS_STORAGE.match(row[col_disk].replace(",", ""))
            if match:
                disk = int(match.group(1)) * float(match.group(2)) * GIB

        try:
            types[name] = entry(row[col_cpu], mem[:-4], arch, disk)
        except ValueError:
            continue

    return types


def get_azure():
    # Cloudflare rejects the default Python user agent.
    data = fetch_json("https://instances.vantage.sh/azure/instances.json", headers={"User-Agent": "instance-types/1.0"})

    types = {}
    for e in data:
        name = e.get("pretty_name_azure")
        cpu = e.get("vcpu")
        mem = e.get("memory")
        if not name or not cpu or mem is None:
            continue

        # The N series is all accelerators (GPU, FPGA).
        if e.get("GPU") not in (None, "", "0", 0) or name.startswith("N"):
            continue

        arch = "aarch64" if "Arm64" in (e.get("arch") or []) else "x86_64"

        # "size" is the local/temporary disk in GiB.
        disk = (e.get("size") or 0) * GIB

        types["Standard_" + name.replace(" ", "_")] = entry(cpu, mem, arch, disk)

    return types


def get_gce():
    # Cloudflare rejects the default Python user agent.
    resp = fetch("https://gcloud-compute.com/machine-types-regions.csv", headers={"User-Agent": "instance-types/1.0"})
    reader = csv.DictReader(io.TextIOWrapper(resp, encoding="utf-8", newline=""))

    types = {}
    for row in reader:
        name = row.get("name")
        if not name or name in types:
            continue

        try:
            if float(row.get("acceleratorCount") or 0) > 0:
                continue

            arch = "aarch64" if float(row.get("arm") or 0) > 0 else "x86_64"

            # Local SSDs are optional add-ons on GCE, no disk is inherent to the type.
            types[name] = entry(row["vCpus"], row["memoryGB"], arch)
        except (KeyError, ValueError):
            continue

    return types


def get_digitalocean():
    # The official /v2/sizes endpoint only returns the sizes available to the
    # authenticating account, this mirror provides the full list without auth.
    data = fetch_json("https://slugs.do-api.dev/api/sizes")

    types = {}
    for size in data["sizes"]:
        if size["slug"].startswith("gpu-"):
            continue

        types[size["slug"]] = entry(size["vcpus"], size["memory"] / 1024, disk=size["disk"] * GIB)

    return types


def get_hetzner():
    # Hetzner has no public endpoint, a (read-only) API token is required.
    token = os.environ.get("HCLOUD_TOKEN")
    if not token:
        print("hetzner: skipped (HCLOUD_TOKEN not set)")
        return None

    headers = {"Authorization": "Bearer " + token}
    types = {}
    page = 1
    while page:
        data = fetch_json(f"https://api.hetzner.cloud/v1/server_types?per_page=50&page={page}", headers=headers)
        for st in data["server_types"]:
            arch = "aarch64" if st.get("architecture") == "arm" else "x86_64"
            types[st["name"]] = entry(st["cores"], st["memory"], arch, st["disk"] * GIB)

        page = data.get("meta", {}).get("pagination", {}).get("next_page")

    return types


def get_linode():
    data = fetch_json("https://api.linode.com/v4/linode/types")

    types = {}
    for e in data["data"]:
        if e.get("class") in ("gpu", "accelerated"):
            continue

        types[e["id"]] = entry(e["vcpus"], e["memory"] / 1024, disk=e["disk"] * 1024 ** 2)

    return types


def get_ovh():
    data = fetch_json("https://api.ovh.com/1.0/order/catalog/public/cloud?ovhSubsidiary=IE")

    types = {}
    for addon in data.get("addons", []):
        if addon.get("product") != "publiccloud-instance":
            continue

        tech = (addon.get("blobs") or {}).get("technical") or {}
        cpu = (tech.get("cpu") or {}).get("cores")
        mem = (tech.get("memory") or {}).get("size")
        name = addon.get("invoiceName")
        if not name or not cpu or not mem:
            continue

        if tech.get("gpu"):
            continue

        # Windows flavors are just licensed variants of the same hardware.
        if name.startswith("win-"):
            continue

        disk = 0
        for d in (tech.get("storage") or {}).get("disks") or []:
            disk += d.get("number", 1) * (d.get("capacity") or 0) * GIB

        types[name] = entry(cpu, mem, disk=disk)

    return types


def get_scaleway():
    zones = (
        "fr-par-1", "fr-par-2", "fr-par-3",
        "nl-ams-1", "nl-ams-2", "nl-ams-3",
        "pl-waw-1", "pl-waw-2", "pl-waw-3",
    )

    arches = {"x86_64": "x86_64", "arm64": "aarch64", "arm": "armv7l"}

    types = {}
    for zone in zones:
        data = fetch_json(f"https://api.scaleway.com/instance/v1/zones/{zone}/products/servers")
        for name, server in data["servers"].items():
            if name in types or server.get("gpu"):
                continue

            # max_size is the total local volume space in bytes, 0 when block storage only.
            disk = (server.get("volumes_constraint") or {}).get("max_size") or 0

            types[name] = entry(server["ncpus"], server["ram"] / GIB,
                                arches.get(server.get("arch"), "x86_64"), disk)

    return types


def get_vultr():
    types = {}
    for endpoint, key in (("plans", "plans"), ("plans-metal", "plans_metal")):
        url = f"https://api.vultr.com/v2/{endpoint}?per_page=500"
        while url:
            data = fetch_json(url)
            for plan in data[key]:
                if plan.get("gpu_vram_gb") or plan["id"].startswith("vcg-"):
                    continue

                cpu = plan.get("vcpu_count", plan.get("cpu_count"))
                disk = plan.get("disk", 0) * plan.get("disk_count", 1) * GIB
                types[plan["id"]] = entry(cpu, plan["ram"] / 1024, disk=disk)

            cursor = data.get("meta", {}).get("links", {}).get("next")
            url = f"https://api.vultr.com/v2/{endpoint}?per_page=500&cursor={cursor}" if cursor else None

    return types


def format_number(value):
    # Render whole numbers as "2.0" and trim float noise elsewhere.
    if value == int(value):
        return f"{value:.1f}"

    return repr(round(value, 3))


def write_yaml(path, types):
    tmp = path + ".tmp"
    with open(tmp, "w", encoding="utf-8") as fd:
        for name in sorted(types):
            data = types[name]
            fd.write(f"{name}:\n")
            fd.write(f"  arch: {data['arch']}\n")
            fd.write(f"  cpu: {format_number(data['cpu'])}\n")
            if "disk" in data:
                fd.write(f"  disk: {data['disk']}\n")

            fd.write(f"  mem: {format_number(data['mem'])}\n")

    os.rename(tmp, path)


def load_yaml(path):
    types = {}
    name = None
    with open(path, encoding="utf-8") as fd:
        for line in fd:
            if not line.startswith(" ") and line.rstrip().endswith(":"):
                name = line.rstrip()[:-1]
                types[name] = {}
            elif name and ":" in line:
                key, value = line.split(":", 1)
                types[name][key.strip()] = value.strip()

    return types


def write_clouds(output):
    tmp = os.path.join(output, "clouds.yaml.tmp")
    with open(tmp, "w", encoding="utf-8") as fd:
        for cloud, cloud_name in CLOUDS:
            if not os.path.exists(os.path.join(output, cloud + ".yaml")):
                continue

            fd.write(f"{cloud}:\n  name: {cloud_name}\n")

    os.rename(tmp, os.path.join(output, "clouds.yaml"))


def format_table_number(value):
    value = float(value)
    if value == int(value):
        return str(int(value))

    return str(round(value, 2))


def write_index(output, template):
    with open(template, encoding="utf-8") as fd:
        tpl = fd.read()

    available = [(cloud, cloud_name) for cloud, cloud_name in CLOUDS
                 if os.path.exists(os.path.join(output, cloud + ".yaml"))]

    sections = ["<ul>"]
    for cloud, cloud_name in available:
        sections.append(f"<li><a href=\"#{cloud}\">{html.escape(cloud_name)}</a></li>")

    sections.append("</ul>")

    for cloud, cloud_name in available:
        path = os.path.join(output, cloud + ".yaml")
        types = load_yaml(path)
        mtime = time.strftime("%Y-%m-%d %H:%M:%S UTC", time.gmtime(os.path.getmtime(path)))
        has_disk = any("disk" in data for data in types.values())

        sections.append(f"<h2 id=\"{cloud}\">{html.escape(cloud_name)}</h2>")
        sections.append(f"<p><a href=\"{cloud}.yaml\">{cloud}.yaml</a> &mdash; "
                        f"{len(types)} instance types &mdash; updated at {mtime}</p>")

        arches = sorted({data.get("arch", "unknown") for data in types.values()})
        for arch in arches:
            if len(arches) > 1:
                sections.append(f"<h3>{html.escape(arch)}</h3>")

            sections.append("<table class=\"table table-striped\">")
            headers = "<th>Instance type</th><th>CPU</th><th>Memory</th>"
            if has_disk:
                headers += "<th>Disk</th>"

            sections.append(f"<thead><tr>{headers}</tr></thead>")
            sections.append("<tbody>")

            for name in sorted(types):
                data = types[name]
                if data.get("arch", "unknown") != arch:
                    continue

                row = (f"<td>{html.escape(name)}</td>"
                       f"<td>{format_table_number(data['cpu'])}</td>"
                       f"<td>{format_table_number(data['mem'])} GiB</td>")
                if has_disk:
                    disk = f"{format_table_number(int(data['disk']) / GIB)} GiB" if "disk" in data else "-"
                    row += f"<td>{disk}</td>"

                sections.append(f"<tr>{row}</tr>")

            sections.append("</tbody>")
            sections.append("</table>")

    now = time.strftime("%Y-%m-%d %H:%M:%S UTC", time.gmtime())
    content = tpl.replace("@DATA@", "\n".join(sections)).replace("@TIMESTAMP@", now)

    tmp = os.path.join(output, "index.html.tmp")
    with open(tmp, "w", encoding="utf-8") as fd:
        fd.write(content)

    os.rename(tmp, os.path.join(output, "index.html"))


def main():
    parser = argparse.ArgumentParser(description="Regenerate instance-types YAML files.")
    parser.add_argument("--output", default=".", help="Output directory (default: current directory)")
    parser.add_argument("--template", default=None, help="index.html template (default: .index.tpl in output directory)")
    args = parser.parse_args()

    os.makedirs(args.output, exist_ok=True)

    template = args.template
    if template is None:
        for candidate in (os.path.join(args.output, ".index.tpl"),
                          os.path.join(os.path.dirname(os.path.abspath(__file__)), ".index.tpl")):
            if os.path.exists(candidate):
                template = candidate
                break

    failed = []
    for cloud, _ in CLOUDS:
        try:
            types = globals()[f"get_{cloud}"]()
            if types is None:
                continue

            if len(types) < MIN_COUNT[cloud]:
                raise RuntimeError(f"only {len(types)} types returned (expected >= {MIN_COUNT[cloud]})")

            write_yaml(os.path.join(args.output, cloud + ".yaml"), types)
            print(f"{cloud}: {len(types)} types")
        except Exception as exc:
            print(f"{cloud}: FAILED: {exc}", file=sys.stderr)
            failed.append(cloud)

    write_clouds(args.output)

    if template:
        write_index(args.output, template)
    else:
        print("index.html: skipped (no .index.tpl found)")

    return 1 if failed else 0


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