#!/usr/bin/env python3
"""Project Euler Problem 1009: matching digits of n and 2n in two bases."""

from __future__ import annotations

import sys

MAX_BASE = 20


def largest_value(a: int, b: int) -> int:
    """Return F(a, b), using exact integer arithmetic throughout."""
    if b >= 3 * a:
        return 0
    if b == 2 * a:
        return a * (a - 1)

    weights: list[int] = []
    power_a = power_b = 1
    capacity = 0

    # The first k with b**k >= 2*a**k bounds the highest digit position.
    while power_b < 2 * power_a:
        weight = 2 * power_a - power_b
        capacity += (a - 1) * weight
        weights.append(weight)
        power_a *= a
        power_b *= b

    leading_weight = power_b - 2 * power_a
    value = min(a - 1, capacity // leading_weight)
    remaining = value * leading_weight

    # w_i <= 1 + (a-1)*sum(w_j, j<i) makes this digitwise greedy exact.
    for weight in reversed(weights):
        digit = min(a - 1, remaining // weight)
        remaining -= digit * weight
        value = value * a + digit
    return value


def sum_for_base(a: int) -> int:
    return sum(largest_value(a, b) for b in range(a + 1, 3 * a))


def same_digits(n: int, doubled: int, a: int, b: int) -> bool:
    while n != 0 or doubled != 0:
        if n % a != doubled % b:
            return False
        n //= a
        doubled //= b
    return True


def brute_force(a: int, b: int) -> int:
    power_a = power_b = 1
    while power_b < 2 * power_a:
        power_a *= a
        power_b *= b

    result = 0
    for n in range(1, power_a * a):
        if same_digits(n, 2 * n, a, b):
            result = n
    return result


def require(condition: bool, description: str) -> None:
    if not condition:
        raise RuntimeError("Check failed: " + description)


def run_tests() -> None:
    require(largest_value(3, 4) == 53, "F(3,4)")
    require(largest_value(9, 10) == 8_152_650, "F(9,10)")
    require(sum_for_base(3) == 72, "G(3)")
    for a in range(2, 8):
        for b in range(a + 1, 3 * a + 3):
            require(
                largest_value(a, b) == brute_force(a, b),
                f"exhaustive comparison for a={a}, b={b}",
            )
    for a in range(2, MAX_BASE + 1):
        for b in range(a + 1, 3 * a):
            value = largest_value(a, b)
            require(
                same_digits(value, 2 * value, a, b),
                f"matching base representations for a={a}, b={b}",
            )
    print("All checks passed.")


def main() -> None:
    if sys.argv[1:] == ["--self-test"]:
        run_tests()
        return
    if len(sys.argv) != 1:
        raise ValueError("Usage: Euler1009.py [--self-test]")
    print(sum(sum_for_base(a) for a in range(2, MAX_BASE + 1)))


if __name__ == "__main__":
    try:
        main()
    except (RuntimeError, ValueError) as error:
        print(error, file=sys.stderr)
        sys.exit(1)
