#!/usr/bin/env python3
"""Reproduce the update-path scenarios used by the accompanying whitepaper.

The model is intentionally narrow. It estimates a resource envelope for moving
one client from an existing build to a launchable state. It does not emulate the
Steam client or estimate the experience of the Steam population.

The scenario file separates documented mechanics from synthetic assumptions.
Each run samples the declared bounded ranges with Python's triangular
distribution. The fixed seed makes the generated CSV files reproducible.

Standard library only. No network access.

Usage:
    python simulate_update.py
    python simulate_update.py --scenario-file update-scenarios.json
    python simulate_update.py --self-test
"""

from __future__ import annotations

import argparse
import csv
import hashlib
import json
import math
import platform
import random
import sys
from dataclasses import dataclass
from pathlib import Path
from typing import Any


MODEL_VERSION = "1.0.0"
PYTHON_VERSION = platform.python_version()
SCENARIO_FILENAME = "update-scenarios.json"
RUNS_FILENAME = "update-simulation-runs.csv"
SUMMARY_FILENAME = "update-simulation-summary.csv"
BUNDLE_URL = (
    "https://jasondoyle.ie/whitepapers/"
    "the-install-button-hides-a-distributed-system/"
    "update-simulator-bundle.zip"
)
MIB = 1024 * 1024
VALVE_EXAMPLE_PACK_MIB = 25_000_000_000 / MIB


class ModelError(Exception):
    """Raised when scenario data violates a model invariant."""


@dataclass(frozen=True)
class Range:
    low: float
    mode: float
    high: float

    @classmethod
    def from_value(cls, value: Any, field: str) -> "Range":
        if type(value) in (int, float):
            number = float(value)
            if not math.isfinite(number):
                raise ModelError(f"{field} must be finite")
            return cls(number, number, number)
        if not isinstance(value, dict):
            raise ModelError(f"{field} must be a number or range object")
        raw_values = [value.get(name) for name in ("low", "mode", "high")]
        if any(type(item) not in (int, float) for item in raw_values):
            raise ModelError(
                f"{field} must contain numeric low, mode and high values"
            )
        try:
            result = cls(
                float(value["low"]),
                float(value["mode"]),
                float(value["high"]),
            )
        except (KeyError, TypeError, ValueError) as error:
            raise ModelError(
                f"{field} must contain numeric low, mode and high values"
            ) from error
        if not all(math.isfinite(item) for item in result.__dict__.values()):
            raise ModelError(f"{field} must contain finite values")
        if not result.low <= result.mode <= result.high:
            raise ModelError(f"{field} must satisfy low <= mode <= high")
        return result

    def sample(self, rng: random.Random) -> float:
        if self.low == self.high:
            return self.low
        return rng.triangular(self.low, self.high, self.mode)


def script_directory() -> Path:
    return Path(__file__).resolve().parent


def scenario_path() -> Path:
    return script_directory() / SCENARIO_FILENAME


def sha256_of_file(path: Path) -> str:
    digest = hashlib.sha256()
    with path.open("rb") as handle:
        for block in iter(lambda: handle.read(65536), b""):
            digest.update(block)
    return digest.hexdigest()


def load_config(path: Path) -> dict[str, Any]:
    try:
        config = json.loads(
            path.read_text(encoding="ascii"),
            parse_constant=lambda value: (_ for _ in ()).throw(
                ValueError(f"non-finite JSON number: {value}")
            ),
        )
    except FileNotFoundError as error:
        raise ModelError(
            f"scenario file not found: {path}\n"
            f"Download and extract the simulator bundle: {BUNDLE_URL}\n"
            f"Alternatively, place {SCENARIO_FILENAME} beside this script or "
            "pass --scenario-file PATH."
        ) from error
    except (json.JSONDecodeError, ValueError) as error:
        raise ModelError(f"invalid JSON in {path}: {error}") from error

    if type(config.get("schema_version")) is not int:
        raise ModelError("schema_version must be an integer")
    if config["schema_version"] != 1:
        raise ModelError("schema_version must be 1")
    if type(config.get("seed")) is not int:
        raise ModelError("seed must be an integer")
    if type(config.get("runs_per_scenario")) is not int:
        raise ModelError("runs_per_scenario must be an integer")
    if config["runs_per_scenario"] < 1:
        raise ModelError("runs_per_scenario must be positive")
    if not isinstance(config.get("scenarios"), list) or not config["scenarios"]:
        raise ModelError("scenarios must be a non-empty array")

    identifiers = [item.get("id") for item in config["scenarios"]]
    if any(not isinstance(identifier, str) or not identifier for identifier in identifiers):
        raise ModelError("every scenario needs a non-empty string id")
    if len(identifiers) != len(set(identifiers)):
        raise ModelError("scenario ids must be unique")
    return config


def require_non_negative(value: float, field: str) -> float:
    if not math.isfinite(value):
        raise ModelError(f"{field} must be finite")
    if value < 0:
        raise ModelError(f"{field} must be non-negative")
    return value


def require_positive(value: float, field: str) -> float:
    if value <= 0:
        raise ModelError(f"{field} must be positive")
    return value


def validate_scenario(scenario: dict[str, Any]) -> None:
    required_numbers = [
        "logical_change_mib",
        "compressed_download_mib",
        "changed_uncompressed_mib",
        "touched_file_mib",
        "free_space_mib",
        "chunk_mib",
        "failure_probability_per_attempt",
        "mean_retry_backoff_seconds",
        "deadline_seconds",
    ]
    for field in required_numbers:
        if field not in scenario:
            raise ModelError(f"{scenario['id']}: missing {field}")
        if type(scenario[field]) not in (int, float):
            raise ModelError(f"{scenario['id']}.{field} must be a number")
        require_non_negative(
            float(scenario[field]),
            f"{scenario['id']}.{field}",
        )

    for field in [
        "network_mbps",
        "disk_read_mib_s",
        "disk_write_mib_s",
        "decompress_mib_s",
        "verify_mib_s",
    ]:
        value = Range.from_value(scenario.get(field), f"{scenario['id']}.{field}")
        require_positive(value.low, f"{scenario['id']}.{field}.low")

    for field in [
        "metadata_seconds",
        "commit_seconds",
        "post_install_seconds",
    ]:
        value = Range.from_value(scenario.get(field), f"{scenario['id']}.{field}")
        require_non_negative(value.low, f"{scenario['id']}.{field}.low")

    overlap = Range.from_value(
        scenario.get("overlap_fraction"),
        f"{scenario['id']}.overlap_fraction",
    )
    if overlap.low < 0 or overlap.high > 1:
        raise ModelError(f"{scenario['id']}.overlap_fraction must be within 0..1")

    probability = float(scenario["failure_probability_per_attempt"])
    if probability >= 1:
        raise ModelError(
            f"{scenario['id']}.failure_probability_per_attempt must be below 1"
        )

    if float(scenario["chunk_mib"]) == 0:
        raise ModelError(f"{scenario['id']}.chunk_mib must be positive")
    if (
        float(scenario["changed_uncompressed_mib"]) > 0
        and float(scenario["compressed_download_mib"]) == 0
    ):
        raise ModelError(
            f"{scenario['id']}: changed content requires a positive download"
        )
    if float(scenario["compressed_download_mib"]) > float(
        scenario["changed_uncompressed_mib"]
    ):
        raise ModelError(
            f"{scenario['id']}: compressed download cannot exceed changed "
            "uncompressed bytes in this model"
        )
    if float(scenario["changed_uncompressed_mib"]) > float(
        scenario["touched_file_mib"]
    ):
        raise ModelError(
            f"{scenario['id']}: changed bytes cannot exceed touched file bytes"
        )


def scenario_seed(global_seed: int, identifier: str) -> int:
    material = f"{global_seed}:{identifier}".encode("ascii")
    return int.from_bytes(hashlib.sha256(material).digest()[:8], "big")


def sample_retry_count(
    chunk_count: int,
    failure_probability: float,
    rng: random.Random,
) -> int:
    """Sample total retries for independent geometric chunk attempts.

    Exact sampling is used for small jobs. A normal approximation to the sum of
    geometric variables is used for large jobs. For success probability q, the
    retry count has mean n*p/q and variance n*p/q^2.
    """

    if chunk_count == 0 or failure_probability == 0:
        return 0

    if chunk_count <= 4096:
        retries = 0
        for _ in range(chunk_count):
            while rng.random() < failure_probability:
                retries += 1
        return retries

    success_probability = 1 - failure_probability
    mean = chunk_count * failure_probability / success_probability
    variance = (
        chunk_count
        * failure_probability
        / (success_probability * success_probability)
    )
    sampled = round(rng.gauss(mean, math.sqrt(variance)))
    return max(0, sampled)


def percentile(values: list[float], percent: float) -> float | None:
    if not values:
        return None
    ordered = sorted(values)
    rank = max(1, math.ceil((percent / 100) * len(ordered)))
    return ordered[min(rank, len(ordered)) - 1]


def safe_ratio(numerator: float, denominator: float) -> float | None:
    if denominator == 0:
        return None
    return numerator / denominator


def fmt(value: float | int | str | None) -> str:
    if value is None:
        return ""
    if isinstance(value, str):
        return value
    if isinstance(value, int):
        return str(value)
    return f"{value:.6f}"


def simulate_run(
    scenario: dict[str, Any],
    run_number: int,
    rng: random.Random,
    scenario_digest: str,
) -> dict[str, Any]:
    compressed_download_mib = float(scenario["compressed_download_mib"])
    changed_uncompressed_mib = float(scenario["changed_uncompressed_mib"])
    touched_file_mib = float(scenario["touched_file_mib"])
    reused_mib = max(0.0, touched_file_mib - changed_uncompressed_mib)
    chunk_mib = float(scenario["chunk_mib"])
    chunk_count = (
        math.ceil(changed_uncompressed_mib / chunk_mib)
        if compressed_download_mib
        else 0
    )

    retry_count = sample_retry_count(
        chunk_count,
        float(scenario["failure_probability_per_attempt"]),
        rng,
    )
    retry_mib = (
        compressed_download_mib * retry_count / chunk_count if chunk_count else 0
    )
    transferred_mib = compressed_download_mib + retry_mib

    network_mbps = Range.from_value(
        scenario["network_mbps"], "network_mbps"
    ).sample(rng)
    disk_read_mib_s = Range.from_value(
        scenario["disk_read_mib_s"], "disk_read_mib_s"
    ).sample(rng)
    disk_write_mib_s = Range.from_value(
        scenario["disk_write_mib_s"], "disk_write_mib_s"
    ).sample(rng)
    decompress_mib_s = Range.from_value(
        scenario["decompress_mib_s"], "decompress_mib_s"
    ).sample(rng)
    verify_mib_s = Range.from_value(
        scenario["verify_mib_s"], "verify_mib_s"
    ).sample(rng)
    metadata_seconds = Range.from_value(
        scenario["metadata_seconds"], "metadata_seconds"
    ).sample(rng)
    commit_seconds = Range.from_value(
        scenario["commit_seconds"], "commit_seconds"
    ).sample(rng)
    post_install_seconds = Range.from_value(
        scenario["post_install_seconds"], "post_install_seconds"
    ).sample(rng)
    overlap_fraction = Range.from_value(
        scenario["overlap_fraction"], "overlap_fraction"
    ).sample(rng)

    network_seconds = (transferred_mib * 8.388608) / network_mbps
    network_seconds += retry_count * float(
        scenario["mean_retry_backoff_seconds"]
    )
    cpu_seconds = (
        changed_uncompressed_mib / decompress_mib_s
        + changed_uncompressed_mib / verify_mib_s
    )
    disk_seconds = reused_mib / disk_read_mib_s
    disk_seconds += touched_file_mib / disk_write_mib_s

    resources = {
        "network": network_seconds,
        "cpu": cpu_seconds,
        "disk": disk_seconds,
    }
    bottleneck = max(resources, key=resources.get)
    pipeline_lower_seconds = resources[bottleneck]
    pipeline_serial_seconds = sum(resources.values())
    pipeline_modelled_seconds = pipeline_lower_seconds + (
        1 - overlap_fraction
    ) * (pipeline_serial_seconds - pipeline_lower_seconds)

    extra_space_mib = touched_file_mib
    blocked_for_space = float(scenario["free_space_mib"]) < extra_space_mib
    launch_seconds = None
    within_deadline = False
    status = "blocked_free_space" if blocked_for_space else "launchable"
    if not blocked_for_space:
        launch_seconds = (
            metadata_seconds
            + pipeline_modelled_seconds
            + commit_seconds
            + post_install_seconds
        )
        within_deadline = launch_seconds <= float(scenario["deadline_seconds"])

    logical_change_mib = float(scenario["logical_change_mib"])
    patch_amplification = safe_ratio(
        compressed_download_mib,
        logical_change_mib,
    )
    local_write_amplification = safe_ratio(
        touched_file_mib,
        compressed_download_mib,
    )

    return {
        "model_version": MODEL_VERSION,
        "python_version": PYTHON_VERSION,
        "scenario_sha256": scenario_digest,
        "scenario_id": scenario["id"],
        "run": run_number,
        "status": status,
        "within_deadline": int(within_deadline),
        "deadline_seconds": float(scenario["deadline_seconds"]),
        "launch_seconds": launch_seconds,
        "metadata_seconds": metadata_seconds,
        "network_seconds": network_seconds,
        "cpu_seconds": cpu_seconds,
        "disk_seconds": disk_seconds,
        "commit_seconds": commit_seconds,
        "post_install_seconds": post_install_seconds,
        "pipeline_lower_seconds": pipeline_lower_seconds,
        "pipeline_serial_seconds": pipeline_serial_seconds,
        "pipeline_modelled_seconds": pipeline_modelled_seconds,
        "overlap_fraction": overlap_fraction,
        "bottleneck": bottleneck,
        "network_mbps": network_mbps,
        "disk_read_mib_s": disk_read_mib_s,
        "disk_write_mib_s": disk_write_mib_s,
        "decompress_mib_s": decompress_mib_s,
        "verify_mib_s": verify_mib_s,
        "chunk_count": chunk_count,
        "retry_count": retry_count,
        "download_mib": compressed_download_mib,
        "retry_mib": retry_mib,
        "transferred_mib": transferred_mib,
        "changed_uncompressed_mib": changed_uncompressed_mib,
        "reused_read_mib": reused_mib,
        "written_mib": touched_file_mib,
        "extra_space_mib": extra_space_mib,
        "free_space_mib": float(scenario["free_space_mib"]),
        "patch_amplification": patch_amplification,
        "local_write_amplification": local_write_amplification,
    }


RUN_FIELDS = [
    "model_version",
    "python_version",
    "scenario_sha256",
    "scenario_id",
    "run",
    "status",
    "within_deadline",
    "deadline_seconds",
    "launch_seconds",
    "metadata_seconds",
    "network_seconds",
    "cpu_seconds",
    "disk_seconds",
    "commit_seconds",
    "post_install_seconds",
    "pipeline_lower_seconds",
    "pipeline_serial_seconds",
    "pipeline_modelled_seconds",
    "overlap_fraction",
    "bottleneck",
    "network_mbps",
    "disk_read_mib_s",
    "disk_write_mib_s",
    "decompress_mib_s",
    "verify_mib_s",
    "chunk_count",
    "retry_count",
    "download_mib",
    "retry_mib",
    "transferred_mib",
    "changed_uncompressed_mib",
    "reused_read_mib",
    "written_mib",
    "extra_space_mib",
    "free_space_mib",
    "patch_amplification",
    "local_write_amplification",
]


def summarize(
    scenario: dict[str, Any],
    rows: list[dict[str, Any]],
    scenario_digest: str,
) -> dict[str, Any]:
    launchable = [row for row in rows if row["status"] == "launchable"]
    launch_seconds = [float(row["launch_seconds"]) for row in launchable]
    within_deadline = sum(int(row["within_deadline"]) for row in rows)
    blocked = len(rows) - len(launchable)
    bottlenecks = {
        resource: sum(row["bottleneck"] == resource for row in launchable)
        for resource in ["network", "cpu", "disk"]
    }

    def values(field: str) -> list[float]:
        return [float(row[field]) for row in launchable]

    return {
        "model_version": MODEL_VERSION,
        "python_version": PYTHON_VERSION,
        "scenario_sha256": scenario_digest,
        "scenario_id": scenario["id"],
        "description": scenario["description"],
        "runs": len(rows),
        "launchable_runs": len(launchable),
        "space_block_rate": safe_ratio(blocked, len(rows)),
        "deadline_seconds": float(scenario["deadline_seconds"]),
        "completion_rate_within_deadline": safe_ratio(
            within_deadline,
            len(rows),
        ),
        "launch_p50_seconds": percentile(launch_seconds, 50),
        "launch_p95_seconds": percentile(launch_seconds, 95),
        "launch_p99_seconds": percentile(launch_seconds, 99),
        "network_p50_seconds": percentile(values("network_seconds"), 50),
        "cpu_p50_seconds": percentile(values("cpu_seconds"), 50),
        "disk_p50_seconds": percentile(values("disk_seconds"), 50),
        "pipeline_lower_p50_seconds": percentile(
            values("pipeline_lower_seconds"),
            50,
        ),
        "pipeline_serial_p50_seconds": percentile(
            values("pipeline_serial_seconds"),
            50,
        ),
        "bottleneck_network_rate": safe_ratio(
            bottlenecks["network"],
            len(launchable),
        ),
        "bottleneck_cpu_rate": safe_ratio(
            bottlenecks["cpu"],
            len(launchable),
        ),
        "bottleneck_disk_rate": safe_ratio(
            bottlenecks["disk"],
            len(launchable),
        ),
        "download_mib": float(scenario["compressed_download_mib"]),
        "touched_file_mib": float(scenario["touched_file_mib"]),
        "extra_space_mib": float(scenario["touched_file_mib"]),
        "patch_amplification": safe_ratio(
            float(scenario["compressed_download_mib"]),
            float(scenario["logical_change_mib"]),
        ),
        "local_write_amplification": safe_ratio(
            float(scenario["touched_file_mib"]),
            float(scenario["compressed_download_mib"]),
        ),
    }


SUMMARY_FIELDS = [
    "model_version",
    "python_version",
    "scenario_sha256",
    "scenario_id",
    "description",
    "runs",
    "launchable_runs",
    "space_block_rate",
    "deadline_seconds",
    "completion_rate_within_deadline",
    "launch_p50_seconds",
    "launch_p95_seconds",
    "launch_p99_seconds",
    "network_p50_seconds",
    "cpu_p50_seconds",
    "disk_p50_seconds",
    "pipeline_lower_p50_seconds",
    "pipeline_serial_p50_seconds",
    "bottleneck_network_rate",
    "bottleneck_cpu_rate",
    "bottleneck_disk_rate",
    "download_mib",
    "touched_file_mib",
    "extra_space_mib",
    "patch_amplification",
    "local_write_amplification",
]


def write_csv(path: Path, fields: list[str], rows: list[dict[str, Any]]) -> None:
    with path.open("w", encoding="ascii", newline="") as handle:
        writer = csv.DictWriter(handle, fieldnames=fields, lineterminator="\n")
        writer.writeheader()
        for row in rows:
            writer.writerow({field: fmt(row.get(field)) for field in fields})


def run_model(config: dict[str, Any], digest: str) -> tuple[
    list[dict[str, Any]],
    list[dict[str, Any]],
]:
    run_rows: list[dict[str, Any]] = []
    summary_rows: list[dict[str, Any]] = []

    for scenario in config["scenarios"]:
        validate_scenario(scenario)
        rng = random.Random(scenario_seed(config["seed"], scenario["id"]))
        rows = [
            simulate_run(scenario, run_number, rng, digest)
            for run_number in range(1, config["runs_per_scenario"] + 1)
        ]
        run_rows.extend(rows)
        summary_rows.append(summarize(scenario, rows, digest))

    return run_rows, summary_rows


def self_test() -> None:
    exact_range = Range.from_value(5, "test")
    assert exact_range.sample(random.Random(1)) == 5

    assert sample_retry_count(100, 0, random.Random(1)) == 0

    test_scenario = {
        "id": "self-test",
        "description": "self-test",
        "logical_change_mib": 10 / MIB,
        "compressed_download_mib": 1,
        "changed_uncompressed_mib": 1,
        "touched_file_mib": VALVE_EXAMPLE_PACK_MIB,
        "free_space_mib": 20_000_000_000 / MIB,
        "chunk_mib": 1,
        "failure_probability_per_attempt": 0,
        "mean_retry_backoff_seconds": 0,
        "deadline_seconds": 600,
        "network_mbps": 1000,
        "disk_read_mib_s": 100,
        "disk_write_mib_s": 100,
        "decompress_mib_s": 500,
        "verify_mib_s": 1000,
        "metadata_seconds": 1,
        "commit_seconds": 1,
        "post_install_seconds": 0,
        "overlap_fraction": 1,
    }
    validate_scenario(test_scenario)
    row = simulate_run(
        test_scenario,
        1,
        random.Random(1),
        "test-digest",
    )
    assert row["status"] == "blocked_free_space"
    assert math.isclose(row["patch_amplification"], MIB / 10)
    assert math.isclose(
        row["local_write_amplification"],
        VALVE_EXAMPLE_PACK_MIB,
    )
    assert row["bottleneck"] == "disk"

    test_scenario["free_space_mib"] = 30_000_000_000 / MIB
    row = simulate_run(
        test_scenario,
        1,
        random.Random(1),
        "test-digest",
    )
    assert row["status"] == "launchable"
    assert row["launch_seconds"] is not None
    assert row["launch_seconds"] >= row["pipeline_lower_seconds"]
    assert row["launch_seconds"] <= (
        row["metadata_seconds"]
        + row["pipeline_serial_seconds"]
        + row["commit_seconds"]
    )

    print("Self-test passed")


def main() -> None:
    parser = argparse.ArgumentParser(description=__doc__)
    parser.add_argument(
        "--scenario-file",
        type=Path,
        help=(
            "path to the scenario JSON file; defaults to "
            f"{SCENARIO_FILENAME} beside this script"
        ),
    )
    parser.add_argument(
        "--self-test",
        action="store_true",
        help="run built-in invariant checks and exit",
    )
    args = parser.parse_args()

    if args.self_test:
        self_test()
        return

    path = (
        args.scenario_file.expanduser().resolve()
        if args.scenario_file
        else scenario_path()
    )
    config = load_config(path)
    digest = sha256_of_file(path)
    run_rows, summary_rows = run_model(config, digest)

    write_csv(script_directory() / RUNS_FILENAME, RUN_FIELDS, run_rows)
    write_csv(
        script_directory() / SUMMARY_FILENAME,
        SUMMARY_FIELDS,
        summary_rows,
    )

    print(
        f"Wrote {RUNS_FILENAME} with {len(run_rows)} runs and "
        f"{SUMMARY_FILENAME} with {len(summary_rows)} scenarios"
    )
    print(f"Python version: {PYTHON_VERSION}")
    print(f"Scenario SHA-256: {digest}")


if __name__ == "__main__":
    try:
        main()
    except ModelError as error:
        print(f"error: {error}", file=sys.stderr)
        sys.exit(1)
