#!/usr/bin/env python3
"""Finite periodic TASEP: exact checks, uniformization and independent paths.

NumPy plus the Python standard library. --check never writes files.
"""

import argparse
from fractions import Fraction
import itertools
import json
import math
from pathlib import Path
import platform

import numpy as np

HERE = Path(__file__).resolve().parent


def build_generator(sites, particles):
    """Integer generator in units r, built by occupied-site bit masks."""
    configurations = list(itertools.combinations(range(1, sites + 1), particles))
    masks = [sum(1 << (site - 1) for site in occupied) for occupied in configurations]
    indices = {mask: index for index, mask in enumerate(masks)}
    q = np.zeros((len(masks), len(masks)), dtype=np.int64)
    for column, mask in enumerate(masks):
        for site in range(sites):
            destination = (site + 1) % sites
            if mask & (1 << site) and not mask & (1 << destination):
                new_mask = mask ^ (1 << site) ^ (1 << destination)
                q[indices[new_mask], column] += 1
                q[column, column] -= 1
    return configurations, q


def exact_characteristic_coefficients(matrix):
    """Faddeev-LeVerrier recurrence in exact Fraction arithmetic, descending."""
    n = len(matrix)
    a = [[Fraction(int(value)) for value in row] for row in matrix]
    b = [[Fraction(i == j) for j in range(n)] for i in range(n)]
    coefficients = [Fraction(1)]
    for degree in range(1, n + 1):
        product = [[sum(a[i][k] * b[k][j] for k in range(n))
                    for j in range(n)] for i in range(n)]
        coefficient = -sum(product[i][i] for i in range(n)) / degree
        coefficients.append(coefficient)
        b = [[product[i][j] + (coefficient if i == j else 0)
              for j in range(n)] for i in range(n)]
    return coefficients


def uniformize(q, initial, time, rate_bound, tolerance):
    """Positive Poisson sum; no eigendecomposition or renormalization.

    After term k the omitted Poisson mass is bounded by
    w_(k+1)/(1-mu/(k+2)), once k+2>mu. This is an analytic truncation
    bound, separate from binary64 roundoff in the calculation.
    """
    transition = np.eye(len(q)) + q / rate_bound
    if np.min(transition) < 0 or np.max(np.abs(transition.sum(axis=0) - 1)) > 1e-14:
        raise AssertionError("Uniformization transition matrix is not stochastic.")
    mu = rate_bound * time
    weight = math.exp(-mu)
    state = np.array(initial, dtype=float)
    probability = weight * state
    term = 0
    while True:
        next_weight = weight * mu / (term + 1)
        ratio_bound = mu / (term + 2)
        tail_bound = next_weight / (1 - ratio_bound) if ratio_bound < 1 else math.inf
        if tail_bound <= tolerance:
            break
        term += 1
        weight = next_weight
        state = transition @ state
        probability += weight * state
        if term > 10000:
            raise AssertionError("Uniformization did not converge.")
    return probability, {"last_included_poisson_degree": term,
                         "poisson_tail_upper_bound": tail_bound,
                         "probability_mass_deficit": float(1 - probability.sum()),
                         "minimum_probability": float(np.min(probability))}


def sample_trajectories(parameters, configurations):
    """Independent event-driven implementation using explicit particle lists.

    This routine neither reads Q nor calls the bit-mask transition constructor.
    Observation times on one path are correlated; different paths are independent.
    """
    rng = np.random.Generator(np.random.PCG64(parameters["random_seed"]))
    rate = parameters["rate_r"]
    times = np.array(parameters["dimensionless_times"]) / rate
    lookup = {configuration: index for index, configuration in enumerate(configurations)}
    counts = np.zeros((len(times), len(configurations)), dtype=np.int64)
    for _ in range(parameters["trajectory_count"]):
        occupied = list(parameters["initial_occupied_sites"])
        clock = 0.0
        observation = 0
        while observation < len(times):
            moves = [(particle, site % parameters["sites"] + 1)
                     for particle, site in enumerate(occupied)
                     if site % parameters["sites"] + 1 not in occupied]
            event_time = clock + rng.exponential(1 / (rate * len(moves)))
            while observation < len(times) and times[observation] <= event_time:
                counts[observation, lookup[tuple(occupied)]] += 1
                observation += 1
            if observation == len(times):
                break
            particle, destination = moves[int(rng.integers(len(moves)))]
            occupied[particle] = destination
            occupied.sort()
            clock = event_time
    return counts


def wilson_interval(count, samples):
    """Approximate 95% pointwise binomial interval, not simultaneous coverage."""
    z = 1.959963984540054
    frequency = count / samples
    denominator = 1 + z * z / samples
    center = (frequency + z * z / (2 * samples)) / denominator
    radius = z * math.sqrt(frequency * (1 - frequency) / samples + z * z / (4 * samples**2)) / denominator
    return [max(0.0, center - radius), min(1.0, center + radius)]


def run(parameters):
    if parameters["sites"] != 4 or parameters["particles"] != 2:
        raise ValueError("This reproduction is explicitly N=4, M=2.")
    rate = parameters["rate_r"]
    times = parameters["dimensionless_times"]
    samples = parameters["trajectory_count"]
    if not math.isfinite(rate) or rate <= 0 or not times or any(t <= 0 for t in times) or sorted(times) != times:
        raise ValueError("Use positive rate and positive increasing dimensionless times.")
    if samples < 1 or parameters["random_generator"] != "PCG64":
        raise ValueError("Use at least one trajectory and the specified PCG64 generator.")
    configurations, unit_q = build_generator(4, 2)
    # Separately specified from the six-state transition graph, not generated.
    expected_q = np.array([[-1, 0, 0, 0, 1, 0], [1, -2, 0, 0, 0, 1],
                           [0, 1, -1, 0, 0, 0], [0, 1, 0, -1, 0, 0],
                           [0, 0, 1, 1, -2, 0], [0, 0, 0, 0, 1, -1]])
    if not np.array_equal(unit_q, expected_q):
        raise AssertionError("Allowed-jump construction disagrees with the transition graph.")
    q = rate * unit_q
    stationary = np.ones(6) / 6
    initial = np.zeros(6)
    initial[configurations.index(tuple(parameters["initial_occupied_sites"]))] = 1
    if parameters["initial_occupied_sites"] != [1, 2]:
        raise ValueError("The direction-sensitive initial-state check uses occupied sites 1,2.")
    off_diagonal = unit_q - np.diag(np.diag(unit_q))
    if np.min(off_diagonal) < 0 or np.any(unit_q.sum(axis=0)) or np.any(unit_q.sum(axis=1)):
        raise AssertionError("Generator positivity, conservation or uniform stationarity failed.")
    characteristic = exact_characteristic_coefficients(unit_q)
    expected_coefficients = [1, 8, 26, 44, 37, 12, 0]
    if characteristic != [Fraction(value) for value in expected_coefficients]:
        raise AssertionError("Exact characteristic polynomial differs from q(q+1)^2(q+3)(q^2+3q+4).")
    eigenvalues = np.linalg.eigvals(unit_q)
    # Sorted multisets avoid unstable ordering of conjugate or repeated roots.
    spectrum = sorted(eigenvalues, key=lambda value: (round(value.real, 10), round(value.imag, 10)))
    analytic_spectrum = sorted([0, -1, -1, -3, (-3 + 1j * math.sqrt(7)) / 2,
                               (-3 - 1j * math.sqrt(7)) / 2], key=lambda value: (value.real, value.imag))
    spectral_error = max(abs(a - b) for a, b in zip(spectrum, analytic_spectrum))
    z1 = complex(-0.5, math.sqrt(3) / 2)
    z2 = z1.conjugate()
    scattering = -(1 - z2) / (1 - z1)
    bethe = np.array([z1**x * z2**y + scattering * z2**x * z1**y for x, y in configurations])
    mode = np.array([1, -2, 1, 1, -2, 1])
    mode_residual = float(np.linalg.norm(unit_q @ bethe + 3 * bethe) / np.linalg.norm(bethe))
    periodic_residual = max(abs(z1**4 - 1 / scattering), abs(z2**4 - scattering))
    mode_shape_error = float(np.max(np.abs(bethe / bethe[0] - mode)))
    expected_derivative = np.array([-1, 1, 0, 0, 0, 0])
    if not np.array_equal(unit_q @ initial, expected_derivative):
        raise AssertionError("The initial transition 12 to 13 has the wrong direction.")
    transpose_direction_error = float(np.linalg.norm(unit_q.T @ initial - expected_derivative, ord=1))
    if transpose_direction_error != 2:
        raise AssertionError("Transposition negative control did not detect the direction error.")
    current_fractions = [Fraction(sum(site in c and site % 4 + 1 not in c for c in configurations), 6)
                         for site in range(1, 5)]
    if current_fractions != [Fraction(1, 3)] * 4:
        raise AssertionError("Exact per-bond stationary current failed.")
    exit_counts = -np.diag(unit_q)
    jump_transition = np.eye(6) + unit_q / exit_counts[np.newaxis, :]
    jump_stationary = exit_counts / exit_counts.sum()
    recovered_stationary = jump_stationary / exit_counts
    recovered_stationary /= recovered_stationary.sum()
    jump_checks = {"exit_counts": exit_counts.tolist(),
                   "stationary_epoch_weights": jump_stationary.tolist(),
                   "stationarity_max_entry_error": float(np.max(np.abs(jump_transition @ jump_stationary - jump_stationary))),
                   "epoch_vs_time_probability_l1_distance": float(np.sum(np.abs(jump_stationary - stationary))),
                   "recovered_time_probability_max_entry_error": float(np.max(np.abs(recovered_stationary - stationary)))}
    amplitude = parameters["mode_perturbation_amplitude"]
    perturbed = stationary + amplitude * mode
    if np.min(perturbed) < 0:
        raise ValueError("The mode perturbation must be a nonnegative probability distribution.")
    counts = sample_trajectories(parameters, configurations)
    alpha = parameters["statistical_family_failure_probability"]
    if not 0 < alpha < 1:
        raise ValueError("Statistical family failure probability must lie between 0 and 1.")
    family_size = 6 * len(times)
    hoeffding = math.sqrt(math.log(2 * family_size / alpha) / (2 * samples))
    evolution = []
    for index, scaled_time in enumerate(times):
        physical_time = scaled_time / rate
        deterministic, diagnostics = uniformize(q, initial, physical_time, 2 * rate,
                                               parameters["uniformization_tail_tolerance"])
        evolved_mode, mode_diagnostics = uniformize(q, perturbed, physical_time, 2 * rate,
                                                  parameters["uniformization_tail_tolerance"])
        exact_mode = stationary + amplitude * math.exp(-3 * scaled_time) * mode
        analytic_error = float(np.sum(np.abs(evolved_mode - exact_mode)))
        empirical = counts[index] / samples
        standard_errors = np.sqrt(deterministic * (1 - deterministic) / samples)
        probability_error = np.abs(empirical - deterministic)
        if counts[index].sum() != samples or np.max(probability_error) > hoeffding + diagnostics["poisson_tail_upper_bound"] + 2e-12:
            raise AssertionError("Trajectory frequencies fail the simultaneous Hoeffding check.")
        if abs(diagnostics["probability_mass_deficit"]) > diagnostics["poisson_tail_upper_bound"] + 2e-12:
            raise AssertionError("Uniformization lost more mass than truncation plus roundoff permits.")
        if analytic_error > mode_diagnostics["poisson_tail_upper_bound"] + 2e-12:
            raise AssertionError("Uniformization disagrees with independently known mode evolution.")
        evolution.append({"dimensionless_time_rt": scaled_time, "physical_time": physical_time,
                          "deterministic_probabilities": deterministic.tolist(), "uniformization": diagnostics,
                          "trajectory_counts": counts[index].tolist(), "empirical_probabilities": empirical.tolist(),
                          "binomial_standard_errors_from_deterministic_p": standard_errors.tolist(),
                          "wilson_95_pointwise_intervals": [wilson_interval(int(count), samples) for count in counts[index]],
                          "max_absolute_probability_error": float(np.max(probability_error)),
                          "max_absolute_standardized_error": float(np.max(probability_error / standard_errors)),
                          "mode_perturbation_analytic_l1_error": analytic_error})
    tolerance = parameters["equation_tolerance"]
    accepted = [spectral_error, mode_residual, periodic_residual, mode_shape_error,
                jump_checks["stationarity_max_entry_error"], jump_checks["recovered_time_probability_max_entry_error"]]
    if max(accepted) > tolerance:
        raise AssertionError("An independent deterministic identity failed.")
    return {"schema_version": 1,
            "environment": {"python": platform.python_version(), "numpy": np.__version__,
                            "arithmetic": "binary64, integer generator, Fraction characteristic polynomial", "rng": "PCG64"},
            "inputs": parameters, "configuration_basis": [list(c) for c in configurations],
            "generator_in_units_r": unit_q.tolist(),
            "exact_checks": {"column_sums": unit_q.sum(axis=0).tolist(), "uniform_stationarity_numerator": unit_q.sum(axis=1).tolist(),
                             "characteristic_coefficients_descending": [str(c) for c in characteristic],
                             "spectrum_in_units_r_real_imag": [[float(v.real), float(v.imag)] for v in spectrum],
                             "spectrum_max_error_in_units_r": float(spectral_error),
                             "per_bond_stationary_current_over_r": [str(c) for c in current_fractions],
                             "total_stationary_jump_rate_over_r": "4/3", "per_particle_jump_rate_over_r": "2/3"},
            "bethe_mode": {"z1_real_imag": [z1.real, z1.imag], "z2_real_imag": [z2.real, z2.imag],
                           "scattering_amplitude_real_imag": [scattering.real, scattering.imag],
                           "mode_proportional_to": mode.tolist(), "eigenvalue_over_r": -3,
                           "relative_eigenvector_residual_over_r": mode_residual,
                           "periodic_equation_residual": float(periodic_residual), "shape_max_entry_error": mode_shape_error},
            "direction_control": {"correct_initial_derivative_over_r": (unit_q @ initial).tolist(),
                                  "transposed_initial_derivative_over_r": (unit_q.T @ initial).tolist(),
                                  "transposed_direction_l1_error_over_r": transpose_direction_error,
                                  "transposed_mode_residual": float(np.linalg.norm(unit_q.T @ mode + 3 * mode))},
            "jump_epoch_bias": jump_checks,
            "statistical_check": {"independent_trajectories": samples, "probability_comparisons": family_size,
                                  "family_failure_probability_bound": alpha, "hoeffding_absolute_probability_tolerance": hoeffding},
            "evolution": evolution}


def compare_saved(actual, saved, atol, rtol, path="results"):
    if isinstance(actual, dict):
        if actual.keys() != saved.keys():
            raise AssertionError(f"Changed fields at {path}")
        for key in actual:
            if key == "inputs":
                if actual[key] != saved[key]:
                    raise AssertionError("Changed input parameters")
            elif key != "environment":
                compare_saved(actual[key], saved[key], atol, rtol, f"{path}.{key}")
    elif isinstance(actual, list):
        if len(actual) != len(saved):
            raise AssertionError(f"Changed length at {path}")
        for index, (left, right) in enumerate(zip(actual, saved)):
            compare_saved(left, right, atol, rtol, f"{path}[{index}]")
    elif isinstance(actual, int):
        if actual != saved:
            raise AssertionError(f"Changed integer at {path}")
    elif isinstance(actual, float):
        if not math.isfinite(actual) or not math.isclose(actual, saved, abs_tol=atol, rel_tol=rtol):
            raise AssertionError(f"Numerical disagreement at {path}: {actual} != {saved}")
    elif actual != saved:
        raise AssertionError(f"Changed value at {path}")


def main():
    parser = argparse.ArgumentParser(description=__doc__)
    flags = parser.add_mutually_exclusive_group()
    flags.add_argument("--check", action="store_true")
    flags.add_argument("--write", action="store_true")
    args = parser.parse_args()
    parameters = json.loads((HERE / "inputs.json").read_text())
    result = run(parameters)
    if args.check:
        tolerance = parameters["saved_comparison"]
        compare_saved(result, json.loads((HERE / "results.json").read_text()),
                      tolerance["absolute_tolerance"], tolerance["relative_tolerance"])
        print("TASEP: generator, current, spectrum, Bethe mode, uniformization, trajectories, negative controls and saved-result checks passed.")
    elif args.write:
        (HERE / "results.json").write_text(json.dumps(result, indent=2, allow_nan=False) + "\n")
        print("Wrote results.json after all TASEP checks passed.")
    else:
        print(json.dumps(result, indent=2, allow_nan=False))


if __name__ == "__main__":
    main()
