package nu.sutic.ncp;

import java.util.ArrayList;
import java.util.List;
import java.util.Comparator;

/**
 * A natural cubic spline, used for interpolation.
 */
public class NaturalSpline {

    public static record Point(double x, double y) {}

    public final int n;
    public final List<Point> points;
    
    private final double[] a;
    private final double[] b;
    private final double[] c;
    private final double[] d;
    
    private final double[] x;
    private final double[] y;

    private final double[] h;
    private final double[] m;

    public NaturalSpline(List<Point> handles) {

        // Sort by X axis
        points = new ArrayList<>(handles);
        points.sort(Comparator.comparingDouble(p -> p.x()));

        n = points.size();

        // Natural Cubic Spline Interpolation
        x = new double[n];
        y = new double[n];
        for (int i = 0; i < n; i++) {
            x[i] = points.get(i).x();
            y[i] = points.get(i).y();
        }

        h = new double[n - 1];
        for (int i = 0; i < n - 1; i++) {
            h[i] = x[i + 1] - x[i];
        }

        // Setting up tridiagonal linear system for the second derivatives (spline curvatures)
        a = new double[n];
        b = new double[n];
        c = new double[n];
        d = new double[n];

        // Natural boundary conditions: second derivatives at endpoints equal 0
        b[0] = 1.0;
        b[n - 1] = 1.0;

        for (int i = 1; i < n - 1; i++) {
            a[i] = h[i - 1];
            b[i] = 2.0 * (h[i - 1] + h[i]);
            c[i] = h[i];
            d[i] = 6.0 * ((y[i + 1] - y[i]) / h[i] - (y[i] - y[i - 1]) / h[i - 1]);
        }

        // Thomas Algorithm Solver for Tridiagonal Matrix System
        double[] cPrime = new double[n];
        double[] dPrime = new double[n];
        m = new double[n];

        cPrime[0] = c[0] / b[0];
        dPrime[0] = d[0] / b[0];

        for (int i = 1; i < n; i++) {
            double mDiv = b[i] - a[i] * cPrime[i - 1];
            if (i < n - 1) {
                cPrime[i] = c[i] / mDiv;
            }
            dPrime[i] = (d[i] - a[i] * dPrime[i - 1]) / mDiv;
        }

        m[n - 1] = dPrime[n - 1];
        for (int i = n - 2; i >= 0; i--) {
            m[i] = dPrime[i] - cPrime[i] * m[i + 1];
        }
    }

    public Point first() {
        return points.get(0);
    }

    public Point last() {
        return points.get(points.size() - 1);
    }

    public double evaluate(double atX) {
        int intervalIdx = 0;

        // Step forward to locate correct localized boundary domain interval
        while (intervalIdx < n - 1 && atX > x[intervalIdx + 1]) {
            intervalIdx++;
        }

        double xL = x[intervalIdx];
        double xR = x[intervalIdx + 1];
        double hInt = h[intervalIdx];

        // Evaluate Spline Function Formula
        double term1 = m[intervalIdx] * Math.pow(xR - atX, 3) / (6.0 * hInt);
        double term2 = m[intervalIdx + 1] * Math.pow(atX - xL, 3) / (6.0 * hInt);
        double term3 = (y[intervalIdx] - (m[intervalIdx] * hInt * hInt) / 6.0) * (xR - atX) / hInt;
        double term4 = (y[intervalIdx + 1] - (m[intervalIdx + 1] * hInt * hInt) / 6.0) * (atX - xL) / hInt;

        return term1 + term2 + term3 + term4;
    }

}