#!/usr/bin/env python3
"""Reproduce the U07 spectral-numerics capstone with Python 3 and NumPy.

Run: python foundation-spectral-capstone.py
Optional JSON file: python foundation-spectral-capstone.py --output results.json
Only the optional output path is written. No network access is used.
"""
import argparse
import json
import platform
from pathlib import Path

import numpy as np


def chebyshev_matrix(degree):
    """Lobatto nodes ordered from +1 to -1 and their interpolation derivative."""
    x = np.cos(np.pi * np.arange(degree + 1) / degree)
    c = np.ones(degree + 1)
    c[[0, -1]] = 2
    c *= (-1.0) ** np.arange(degree + 1)
    differences = x[:, None] - x[None, :]
    D = (c[:, None] / c[None, :]) / (differences + np.eye(degree + 1))
    D -= np.diag(D.sum(axis=1))
    assert np.max(np.abs(D @ np.ones(degree + 1))) < 1e-10
    return x, D


def fourier_derivative_checks():
    results = []
    for count in (9, 17, 33):
        x = 2 * np.pi * np.arange(count) / count
        values = np.exp(np.sin(x))
        frequencies = np.fft.fftfreq(count) * count
        numerical = np.fft.ifft(1j * frequencies * np.fft.fft(values)).real
        exact = values * np.cos(x)
        error = float(np.max(np.abs(numerical - exact)))
        results.append({"nodes": count, "max_node_error": error})
    assert results[-1]["max_node_error"] < 1e-12
    return results


def chebyshev_boundary_checks():
    results = []
    for degree in (8, 12, 16, 24):
        x, D = chebyshev_matrix(degree)
        D2 = D @ D  # Square the full matrix before selecting interior rows/columns.
        solution = np.zeros(degree + 1)
        solution[1:-1] = np.linalg.solve(-D2[1:-1, 1:-1], -np.exp(x[1:-1]))
        exact = np.exp(x) - np.cosh(1) - x * np.sinh(1)
        error = float(np.max(np.abs(solution - exact)))
        residual = float(np.max(np.abs(-D2[1:-1] @ solution + np.exp(x[1:-1]))))
        results.append({"degree": degree, "max_node_error": error,
                        "max_interior_residual": residual})
    assert results[2]["max_node_error"] < 1e-12
    return results


def aliasing_check():
    K, N, M = 3, 7, 10
    assert N == 2 * K + 1 and M > 3 * K
    x = 2 * np.pi * np.arange(N) / N
    u = np.cos(3 * x)
    original_coefficients = np.fft.fft(u) / N
    aliased_product = np.fft.fft(u * u) / N
    modes = np.rint(np.fft.fftfreq(N) * N).astype(int)

    padded = np.zeros(M, dtype=complex)
    for j, k in enumerate(modes):
        padded[k % M] = original_coefficients[j]
    fine_values = np.fft.ifft(M * padded)
    fine_product = np.fft.fft(fine_values * fine_values) / M
    retained = np.array([fine_product[k % M] for k in modes])
    reconstructed = np.fft.ifft(N * retained).real

    expected_aliased = np.zeros(N)
    expected_aliased[0] = 0.5
    expected_aliased[1] = expected_aliased[-1] = 0.25
    assert np.max(np.abs(aliased_product - expected_aliased)) < 1e-12
    error = float(np.max(np.abs(reconstructed - 0.5)))
    assert error < 1e-12
    return {"original_nodes": N, "padded_nodes": M, "retained_band": [-K, K],
            "aliased_mode_1_coefficient": float(aliased_product[1].real),
            "dealiased_error_against_constant_half": error}


def main():
    parser = argparse.ArgumentParser(description=__doc__)
    parser.add_argument("--output", type=Path)
    args = parser.parse_args()
    results = {"python": platform.python_version(), "numpy": np.__version__,
               "arithmetic": "IEEE binary64 NumPy arrays",
               "error_metric": "maximum absolute error over the computation nodes",
               "fourier_derivative": fourier_derivative_checks(),
               "chebyshev_boundary_problem": chebyshev_boundary_checks(),
               "quadratic_aliasing": aliasing_check(), "assertions": "passed"}
    text = json.dumps(results, indent=2, ensure_ascii=False) + "\n"
    if args.output:
        args.output.write_text(text, encoding="utf-8")
    print(text, end="")


if __name__ == "__main__":
    main()
