#!/usr/bin/env python3
"""Project Euler Problem 1006 - Fibonacci Subwords."""

from __future__ import annotations

import bisect
import sys
from dataclasses import dataclass

MOD = 101_001_001
BASE = 10
TARGET = 1_000_000_000_000_000_000


def normalize(value: int) -> int:
    return value % MOD


def mod_pow(base: int, exponent: int) -> int:
    return pow(normalize(base), exponent, MOD)


def extended_gcd(a: int, b: int) -> tuple[int, int, int]:
    if b == 0:
        return a, 1, 0
    gcd, next_x, next_y = extended_gcd(b, a % b)
    return gcd, next_y, next_x - (a // b) * next_y


def mod_inverse(value: int) -> int:
    gcd, x, _ = extended_gcd(normalize(value), MOD)
    if gcd != 1:
        raise ValueError("a required modular inverse does not exist")
    return normalize(x)


@dataclass
class Node:
    length: int
    left: int = -1
    right: int = -1
    bit: int = -1


@dataclass
class Summary:
    forward: int = 0
    reverse: int = 0
    pairs: int = 0
    ones: int = 0
    ready: bool = False


class FibonacciSubwords:
    def __init__(self) -> None:
        self.bases = [BASE, mod_inverse(BASE)]
        self.nodes: list[Node] = []
        self.summary_cache: list[list[Summary]] = []
        self.concatenation_cache: dict[tuple[int, int], int] = {}
        self.prefix_cache: dict[tuple[int, int], int] = {}
        self.power_cache: list[dict[int, int]] = [{}, {}]
        self.sum_query_cache: dict[tuple[int, int, int, int, int], int] = {}
        self.difference_query_cache: dict[tuple[int, int, int, int, int], int] = {}

        self.zero = self.add_bit(0)
        self.one = self.add_bit(1)
        self.fibonacci = [1, 2]
        self.standard = [self.zero, self.concatenate(self.zero, self.one)]

        while self.fibonacci[-1] <= TARGET:
            self.fibonacci.append(self.fibonacci[-1] + self.fibonacci[-2])
            self.standard.append(
                self.concatenate(self.standard[-1], self.standard[-2])
            )

    def add_bit(self, bit: int) -> int:
        node_id = len(self.nodes)
        self.nodes.append(Node(length=1, bit=bit))
        self.summary_cache.append([Summary(), Summary()])
        return node_id

    def concatenate(self, left: int, right: int) -> int:
        if left < 0:
            return right
        if right < 0:
            return left

        key = (left, right)
        cached = self.concatenation_cache.get(key)
        if cached is not None:
            return cached

        node_id = len(self.nodes)
        self.nodes.append(
            Node(
                length=self.nodes[left].length + self.nodes[right].length,
                left=left,
                right=right,
            )
        )
        self.summary_cache.append([Summary(), Summary()])
        self.concatenation_cache[key] = node_id
        return node_id

    def prefix(self, standard_index: int, length: int) -> int:
        assert 0 < length <= self.fibonacci[standard_index]
        if length == self.fibonacci[standard_index]:
            return self.standard[standard_index]

        key = (standard_index, length)
        cached = self.prefix_cache.get(key)
        if cached is not None:
            return cached

        assert standard_index > 0
        if length <= self.fibonacci[standard_index - 1]:
            result = self.prefix(standard_index - 1, length)
        else:
            result = self.concatenate(
                self.standard[standard_index - 1],
                self.prefix(
                    standard_index - 2,
                    length - self.fibonacci[standard_index - 1],
                ),
            )

        self.prefix_cache[key] = result
        return result

    def power(self, base_id: int, exponent: int) -> int:
        cache = self.power_cache[base_id]
        cached = cache.get(exponent)
        if cached is not None:
            return cached
        result = mod_pow(self.bases[base_id], exponent)
        cache[exponent] = result
        return result

    def signed_power(self, base_id: int, exponent: int) -> int:
        if exponent >= 0:
            return self.power(base_id, exponent)
        return self.power(1 - base_id, -exponent)

    def summarize(self, node_id: int, base_id: int) -> Summary:
        cached = self.summary_cache[node_id][base_id]
        if cached.ready:
            return cached

        node = self.nodes[node_id]
        if node.bit >= 0:
            result = Summary(
                forward=node.bit,
                reverse=node.bit,
                ones=node.bit,
                ready=True,
            )
            self.summary_cache[node_id][base_id] = result
            return result

        left = self.summarize(node.left, base_id)
        right = self.summarize(node.right, base_id)
        left_length = self.nodes[node.left].length
        right_length = self.nodes[node.right].length
        base = self.bases[base_id]

        result = Summary(
            forward=normalize(
                left.forward + self.power(base_id, left_length) * right.forward
            ),
            reverse=normalize(
                self.power(base_id, right_length) * left.reverse + right.reverse
            ),
            pairs=normalize(
                left.pairs
                + right.pairs
                + base * left.reverse % MOD * right.forward
            ),
            ones=normalize(left.ones + right.ones),
            ready=True,
        )
        self.summary_cache[node_id][base_id] = result
        return result

    def sum_query(
        self, x: int, y: int, low: int, high: int, base_id: int
    ) -> int:
        if x > y:
            x, y = y, x
        key = (x, y, low, high, base_id)
        cached = self.sum_query_cache.get(key)
        if cached is not None:
            return cached

        node_x = self.nodes[x]
        node_y = self.nodes[y]
        maximum = node_x.length + node_y.length - 2

        if high < 0 or low > maximum:
            result = 0
        elif low <= 0 and maximum <= high:
            result = (
                self.summarize(x, base_id).forward
                * self.summarize(y, base_id).forward
                % MOD
            )
        elif node_x.bit >= 0 and node_y.bit >= 0:
            result = node_x.bit * node_y.bit if low <= 0 <= high else 0
        elif node_x.length >= node_y.length and node_x.bit < 0:
            offset = self.nodes[node_x.left].length
            result = normalize(
                self.sum_query(node_x.left, y, low, high, base_id)
                + self.power(base_id, offset)
                * self.sum_query(
                    node_x.right, y, low - offset, high - offset, base_id
                )
            )
        else:
            offset = self.nodes[node_y.left].length
            result = normalize(
                self.sum_query(x, node_y.left, low, high, base_id)
                + self.power(base_id, offset)
                * self.sum_query(
                    x, node_y.right, low - offset, high - offset, base_id
                )
            )

        self.sum_query_cache[key] = result
        return result

    def difference_query(
        self, x: int, y: int, low: int, high: int, base_id: int
    ) -> int:
        key = (x, y, low, high, base_id)
        cached = self.difference_query_cache.get(key)
        if cached is not None:
            return cached

        node_x = self.nodes[x]
        node_y = self.nodes[y]
        minimum = -(node_x.length - 1)
        maximum = node_y.length - 1

        if high < minimum or low > maximum:
            result = 0
        elif low <= minimum and maximum <= high:
            result = (
                self.summarize(x, 1 - base_id).forward
                * self.summarize(y, base_id).forward
                % MOD
            )
        elif node_x.bit >= 0 and node_y.bit >= 0:
            result = node_x.bit * node_y.bit if low <= 0 <= high else 0
        elif node_x.length >= node_y.length and node_x.bit < 0:
            offset = self.nodes[node_x.left].length
            result = normalize(
                self.difference_query(node_x.left, y, low, high, base_id)
                + self.signed_power(base_id, -offset)
                * self.difference_query(
                    node_x.right, y, low + offset, high + offset, base_id
                )
            )
        else:
            offset = self.nodes[node_y.left].length
            result = normalize(
                self.difference_query(x, node_y.left, low, high, base_id)
                + self.power(base_id, offset)
                * self.difference_query(
                    x, node_y.right, low - offset, high - offset, base_id
                )
            )

        self.difference_query_cache[key] = result
        return result

    def duplicated_window_sum(
        self, standard_index: int, k: int, count: int
    ) -> int:
        if count == 0:
            return 0

        initial_word = self.prefix(standard_index, k)
        initial_value = self.summarize(initial_word, 0).reverse
        if count == 1:
            return initial_value * initial_value % MOD

        transition_count = count - 1
        outgoing = self.prefix(standard_index, transition_count)
        outgoing_summary = self.summarize(outgoing, 0)
        window_power = self.power(0, k)
        forward_delta = normalize(
            outgoing_summary.reverse - window_power * outgoing_summary.forward
        )
        reverse_delta = normalize(
            outgoing_summary.forward - window_power * outgoing_summary.reverse
        )

        low_inverse = self.sum_query(
            outgoing, outgoing, 0, transition_count - 2, 1
        )
        high_base = self.sum_query(
            outgoing, outgoing, transition_count, 2 * transition_count - 2, 0
        )
        diagonal_base = self.sum_query(
            outgoing, outgoing, transition_count - 1, transition_count - 1, 0
        )

        diagonal = diagonal_base * self.power(1, transition_count - 1) % MOD
        lower_cross = (
            self.power(0, transition_count - 1) * low_inverse % MOD
        )
        upper_cross = (
            self.power(1, transition_count - 1) * high_base % MOD
        )
        delta_pairs = normalize(
            normalize(1 + window_power * window_power) * outgoing_summary.pairs
            - window_power * normalize(lower_cross + upper_cross)
        )
        delta_squares = normalize(
            normalize(1 + window_power * window_power) * outgoing_summary.ones
            - 2 * window_power % MOD * diagonal
        )

        last_value = normalize(
            self.power(0, transition_count) * initial_value + reverse_delta
        )
        value_delta_sum = normalize(
            initial_value * forward_delta + self.bases[1] * delta_pairs
        )
        numerator = normalize(
            initial_value * initial_value
            - BASE * BASE % MOD * last_value % MOD * last_value
            + 2 * BASE % MOD * value_delta_sum
            + delta_squares
        )
        return numerator * mod_inverse(1 - BASE * BASE) % MOD

    def solve(self, k: int) -> int:
        index = bisect.bisect_right(self.fibonacci, k)
        assert index < len(self.fibonacci)
        cycle_length = self.fibonacci[index]
        duplicate_count = cycle_length - 1 - k
        gap_length = duplicate_count + 1
        word = self.standard[index]

        word_base = self.summarize(word, 0)
        word_inverse = self.summarize(word, 1)
        cycle_power = self.power(0, cycle_length)
        inverse_cycle_power = self.power(1, cycle_length)

        all_correlations_base = normalize(
            word_base.pairs + cycle_power * word_inverse.pairs
        )
        all_correlations_inverse = normalize(
            word_inverse.pairs + inverse_cycle_power * word_base.pairs
        )

        prefix_correlations_base = normalize(
            self.difference_query(word, word, 1, gap_length, 0)
            + cycle_power
            * self.difference_query(
                word,
                word,
                cycle_length - gap_length,
                cycle_length - 1,
                1,
            )
        )
        prefix_correlations_inverse = normalize(
            self.difference_query(word, word, 1, gap_length, 1)
            + inverse_cycle_power
            * self.difference_query(
                word,
                word,
                cycle_length - gap_length,
                cycle_length - 1,
                0,
            )
        )

        used_correlations_base = normalize(
            all_correlations_base
            - cycle_power * prefix_correlations_inverse
        )
        used_correlations_inverse = normalize(
            all_correlations_inverse
            - inverse_cycle_power * prefix_correlations_base
        )

        inverse_geometric_denominator = mod_inverse(BASE * BASE - 1)
        geometric_squares = (
            normalize(self.power(0, 2 * k) - 1)
            * inverse_geometric_denominator
            % MOD
        )
        full_cycle = normalize(
            word_base.ones * geometric_squares
            + 2
            * inverse_geometric_denominator
            % MOD
            * normalize(
                self.power(0, 2 * k) * used_correlations_inverse
                - used_correlations_base
            )
        )

        return normalize(
            full_cycle
            - self.duplicated_window_sum(index, k, duplicate_count)
        )


def brute_force(k: int) -> int:
    older = "0"
    newer = "01"
    while len(newer) < 10 * k + 20:
        older, newer = newer, newer + older

    factors = {
        newer[start : start + k]
        for start in range(len(newer) - k + 1)
    }
    assert len(factors) == k + 1

    result = 0
    for factor in factors:
        value = 0
        for digit in factor:
            value = (BASE * value + ord(digit) - ord("0")) % MOD
        result = (result + value * value) % MOD
    return result


def run_checkpoints() -> None:
    solver = FibonacciSubwords()
    if brute_force(3) != 20_302:
        raise AssertionError("checkpoint failed: Psi(3)")
    for k in range(1, 51):
        if solver.solve(k) != brute_force(k):
            raise AssertionError(
                f"checkpoint failed: brute-force comparison for k={k}"
            )
    if solver.solve(10) != 10_699_667:
        raise AssertionError("checkpoint failed: supplied Psi(10)")
    print("Validation checkpoints passed.", file=sys.stderr)


def main() -> None:
    run_checkpoints()
    solver = FibonacciSubwords()
    print(solver.solve(TARGET))


if __name__ == "__main__":
    main()
