#!/usr/bin/env python3
"""Project Euler Problem 1008 - Functional Inverse."""

from __future__ import annotations

MOD = 1_000_000_007
TARGET = 10_000_000
WANTED_DEGREE = 9  # P_N has a factor x, so x^10 comes from degree 9.
BATCH_SIZE = 32_768


def target_coefficient(n: int, batch_size: int = BATCH_SIZE) -> int:
    """Return the x^10 coefficient in the (n, MOD)-functional inverse.

    If A_k(x) = product(1 - x/j^2, j=1..k), Newton interpolation gives

        P_N(x) / x = sum(A_(k-1)(x) / (k(2k-1)), k=1..N).

    Only coefficients through degree nine are retained.  A batch inversion
    supplies every 1/(k(2k-1)) with one modular exponentiation per batch.
    """

    if n < 1 or 2 * n - 1 >= MOD:
        raise ValueError("this implementation requires 1 <= n and 2n-1 < MOD")
    if batch_size < 1:
        raise ValueError("batch_size must be positive")

    # coefficients[d] is [x^d] A_k for the current k.
    coefficients = [0] * (WANTED_DEGREE + 1)
    coefficients[0] = 1
    interpolation_coefficient = 0
    factorial = 1
    prefix = [1] * (min(batch_size, n) + 1)

    for first in range(1, n + 1, batch_size):
        count = min(batch_size, n - first + 1)

        # Prefix products of v_k = k(2k-1).
        prefix[0] = 1
        for index in range(1, count + 1):
            k = first + index - 1
            value = k * (2 * k - 1) % MOD
            prefix[index] = prefix[index - 1] * value % MOD

        # Recover every v_k^-1 by walking backward from the inverse of the
        # complete product.  Reuse prefix[] to store the resulting inverses.
        suffix_inverse = pow(prefix[count], MOD - 2, MOD)
        for index in range(count, 0, -1):
            k = first + index - 1
            value = k * (2 * k - 1) % MOD
            weight = suffix_inverse * prefix[index - 1] % MOD
            suffix_inverse = suffix_inverse * value % MOD
            prefix[index] = weight

        for index in range(1, count + 1):
            k = first + index - 1
            weight = prefix[index]  # 1/(k(2k-1))
            inverse_k = (2 * k - 1) * weight % MOD
            inverse_square = inverse_k * inverse_k % MOD

            # The summand uses A_(k-1), so accumulate before updating A.
            interpolation_coefficient = (
                interpolation_coefficient
                + weight * coefficients[WANTED_DEGREE]
            ) % MOD
            for degree in range(WANTED_DEGREE, 0, -1):
                coefficients[degree] = (
                    coefficients[degree]
                    - inverse_square * coefficients[degree - 1]
                ) % MOD

            factorial = factorial * k % MOD

    factorial_square = factorial * factorial % MOD
    leading = n * pow(
        (2 * n - 1) * factorial_square % MOD, MOD - 2, MOD
    ) % MOD
    if n % 2 == 0:
        leading = -leading % MOD

    # P_N is already the minimum monic polynomial only when its leading
    # coefficient is one.  Otherwise add the monic vanishing polynomial
    # x*product(x-k^2, k=1..N).
    if leading == 1:
        return interpolation_coefficient

    vanishing_scale = factorial_square if n % 2 == 0 else -factorial_square % MOD
    return (
        interpolation_coefficient
        + vanishing_scale * coefficients[WANTED_DEGREE]
    ) % MOD


def append_root(polynomial: list[int], root: int) -> list[int]:
    """Return polynomial * (x-root), with coefficients modulo MOD."""

    result = [0] * (len(polynomial) + 1)
    for degree, coefficient in enumerate(polynomial):
        result[degree] = (result[degree] - root * coefficient) % MOD
        result[degree + 1] = (result[degree + 1] + coefficient) % MOD
    return result


def interpolate_directly(n: int) -> int:
    """Build the whole small polynomial by Lagrange interpolation.

    This cubic-time reference is intentionally independent of the truncated
    Newton recurrence and is used only by the checkpoints below.
    """

    polynomial = [0] * (n + 2)
    for i in range(1, n + 1):
        basis = [1]
        denominator = 1
        node_i = i * i % MOD
        for j in range(n + 1):
            if j == i:
                continue
            node_j = j * j % MOD
            basis = append_root(basis, node_j)
            denominator = denominator * (node_i - node_j) % MOD

        scale = i * pow(denominator, MOD - 2, MOD) % MOD
        for degree, coefficient in enumerate(basis):
            polynomial[degree] = (
                polynomial[degree] + scale * coefficient
            ) % MOD

    if polynomial[n] != 1:
        vanishing = [1]
        for i in range(n + 1):
            vanishing = append_root(vanishing, i * i % MOD)
        for degree, coefficient in enumerate(vanishing):
            polynomial[degree] = (polynomial[degree] + coefficient) % MOD

    return polynomial[10] if len(polynomial) > 10 else 0


def run_checkpoints() -> None:
    for n in range(1, 41):
        expected = interpolate_directly(n)
        assert target_coefficient(n) == expected, f"interpolation check for n={n}"
        assert target_coefficient(n, 7) == expected, f"batch check for n={n}"


def main() -> None:
    run_checkpoints()
    print(target_coefficient(TARGET))


if __name__ == "__main__":
    main()
