Problem 261: Pivotal Square Sums

View on Project Euler

Project Euler Problem 261 Solution

EulerSolve provides an optimized solution for Project Euler Problem 261, Pivotal Square Sums, with C++, Python, Java, and a step-by-step mathematical explanation.

Problem Summary A positive integer \(k\) is called a square-pivot if there exist integers \(m>0\) and \(n\ge k\) such that $$ (k-m)^2+(k-m+1)^2+\cdots+k^2=(n+1)^2+(n+2)^2+\cdots+(n+m)^2. $$ So the \((m+1)\) consecutive squares ending at \(k\) must equal the \(m\) consecutive squares starting just after \(n\). The task is to sum all distinct square-pivots $$ k\le L, $$ with \(L=10^{10}\) in the full problem. The code does not brute-force all triples \((k,m,n)\); it converts the condition into a Pell-equation pipeline. Mathematical Approach 1. Expanding the Square-Sum Identity Start from $$ \sum_{i=0}^{m}(k-i)^2=\sum_{j=1}^{m}(n+j)^2. $$ Using the standard sum-of-squares formula, the left side becomes $$ (m+1)k^2-km(m+1)+\frac{m(m+1)(2m+1)}{6}, $$ while the right side becomes $$ mn^2+m(m+1)n+\frac{m(m+1)(2m+1)}{6}. $$ The cubic tail cancels immediately, leaving the key Diophantine relation $$ (m+1)k(k-m)=m\,n(n+m+1). $$ This is the real starting point of the implementation. 2. Quadratic Viewpoint and the Brute-Force Check For fixed \(k\) and \(m\), we can treat the previous identity as a quadratic equation in \(n\): $$ n^2+(m+1)n-\frac{(m+1)k(k-m)}{m}=0. $$ Therefore \(n\) is integral if and only if the discriminant $$ \Delta=(m+1)^2+4\frac{(m+1)k(k-m)}{m} $$ is a perfect square and the parity matches....

Detailed mathematical approach

Problem Summary

A positive integer \(k\) is called a square-pivot if there exist integers \(m>0\) and \(n\ge k\) such that

$$ (k-m)^2+(k-m+1)^2+\cdots+k^2=(n+1)^2+(n+2)^2+\cdots+(n+m)^2. $$

So the \((m+1)\) consecutive squares ending at \(k\) must equal the \(m\) consecutive squares starting just after \(n\).

The task is to sum all distinct square-pivots

$$ k\le L, $$

with \(L=10^{10}\) in the full problem. The code does not brute-force all triples \((k,m,n)\); it converts the condition into a Pell-equation pipeline.

Mathematical Approach

1. Expanding the Square-Sum Identity

Start from

$$ \sum_{i=0}^{m}(k-i)^2=\sum_{j=1}^{m}(n+j)^2. $$

Using the standard sum-of-squares formula, the left side becomes

$$ (m+1)k^2-km(m+1)+\frac{m(m+1)(2m+1)}{6}, $$

while the right side becomes

$$ mn^2+m(m+1)n+\frac{m(m+1)(2m+1)}{6}. $$

The cubic tail cancels immediately, leaving the key Diophantine relation

$$ (m+1)k(k-m)=m\,n(n+m+1). $$

This is the real starting point of the implementation.

2. Quadratic Viewpoint and the Brute-Force Check

For fixed \(k\) and \(m\), we can treat the previous identity as a quadratic equation in \(n\):

$$ n^2+(m+1)n-\frac{(m+1)k(k-m)}{m}=0. $$

Therefore \(n\) is integral if and only if the discriminant

$$ \Delta=(m+1)^2+4\frac{(m+1)k(k-m)}{m} $$

is a perfect square and the parity matches.

This is exactly what the checkpoint function brute_pivots tests: it loops over \(k\) and \(m\), computes \(\Delta\), checks whether \(\sqrt{\Delta}\) is integral, and then reconstructs \(n\).

3. A Symmetric Quadratic Form

Introduce the shifted variables

$$ A=2n+m+1,\qquad B=2k-m. $$

Then

$$ n(n+m+1)=\frac{A^2-(m+1)^2}{4}, $$

and

$$ k(k-m)=\frac{B^2-m^2}{4}. $$

Substituting these into the previous identity gives

$$ mA^2-(m+1)B^2=m(m+1). $$

This is much closer to Pell form, because the inhomogeneous right-hand side is exactly the product \(m(m+1)\).

4. Squarefree Decomposition and Pell Reduction

Write

$$ m(m+1)=s q^2, $$

where \(s\) is squarefree and \(q\) is the square part.

Now make the linear substitution

$$ A=(m+1)x+qsy,\qquad B=mx+qsy. $$

A direct expansion shows that

$$ mA^2-(m+1)B^2=s q^2(x^2-sy^2). $$

But the left side must equal \(m(m+1)=s q^2\). After cancelling \(s q^2\), we obtain the Pell equation

$$ x^2-sy^2=1. $$

So for a fixed \(m\), admissible pivots come from solutions of the Pell equation attached to the squarefree part of \(m(m+1)\).

5. Recovering \(k\) and \(n\) from Pell Solutions

Since \(B=2k-m\) and \(A=2n+m+1\), the inverse formulas are

$$ k=\frac{mx+qsy+m}{2}, $$

$$ n=\frac{(m+1)x+qsy-(m+1)}{2}. $$

Also

$$ x=A-B=2(n-k+m)+1. $$

This immediately explains three filters in the code:

1. \(x\) must be odd, because \(2(n-k+m)+1\) is odd,

2. \(n\ge k\) is equivalent to \(x\ge 2m+1\),

3. the numerator of the formula for \(k\) must be even so that \(k\) is integral.

The C++ code implements these as sol.x >= 2m+1, oddness of sol.x, and a final parity check on the numerator.

6. Worked Examples

The official examples fall out naturally from the Pell parametrization.

For \(m=1\), we have

$$ m(m+1)=2=2\cdot 1^2, $$

so \(s=2\) and \(q=1\). The Pell equation is

$$ x^2-2y^2=1. $$

The solution \((x,y)=(3,2)\) gives

$$ k=\frac{1\cdot 3+1\cdot 2\cdot 2+1}{2}=4,\qquad n=\frac{2\cdot 3+1\cdot 2\cdot 2-2}{2}=4. $$

The next Pell solution \((17,12)\) gives

$$ k=\frac{17+24+1}{2}=21,\qquad n=\frac{34+24-2}{2}=28. $$

For \(m=3\),

$$ m(m+1)=12=3\cdot 2^2, $$

so \(s=3\), \(q=2\), and the Pell solution \((x,y)=(7,4)\) gives

$$ k=\frac{3\cdot 7+2\cdot 3\cdot 4+3}{2}=24. $$

For \(m=2\), we get \(s=6\), \(q=1\), and \((x,y)=(49,20)\) gives

$$ k=\frac{2\cdot 49+6\cdot 20+2}{2}=110. $$

These recover the sample pivots \(4\), \(21\), \(24\), and \(110\).

7. Why the Code's Bounds Are Correct

The implementation does not let \(m\) run arbitrarily.

Because \(n\ge k\), we have

$$ n(n+m+1)\ge k(k+m+1). $$

Combining this with

$$ (m+1)k(k-m)=m\,n(n+m+1) $$

gives

$$ (m+1)k(k-m)\ge m k(k+m+1). $$

After dividing by \(k>0\), this simplifies to

$$ k\ge 2m(m+1). $$

Therefore, if \(k\le L\), then necessarily

$$ 2m(m+1)\le L. $$

This yields the code's bound

$$ m\le \left\lfloor\frac{\sqrt{2L+1}-1}{2}\right\rfloor. $$

There is also an \(x\)-bound. Since

$$ k=\frac{mx+qsy+m}{2}\ge \frac{mx+m}{2}, $$

the condition \(k\le L\) implies

$$ x\le \frac{2L-m}{m}. $$

This is exactly the per-entry bound stored as x_max.

8. Why Grouping by Squarefree Part Matters

Different values of \(m\) can have the same squarefree part \(s\) in the factorization of \(m(m+1)\).

Since the Pell equation depends only on \(s\), the code groups all such \(m\) together, solves the Pell equation once for that \(s\), and reuses the entire Pell solution stream for every entry in the group.

This is the main asymptotic improvement over solving a fresh Pell problem for every \(m\).

Different \((m,n)\) pairs can even produce the same pivot \(k\). For example, the pivot

$$ 684 $$

occurs both with \((m,n)=(4,760)\) and with \((m,n)=(18,684)\). That is why the final pivot list must be globally sorted and deduplicated.

How the Code Works

The function build_primes constructs a prime list for factor support. The helper squarefree_and_square_part writes each number as \(s q^2\) with squarefree \(s\).

For each admissible \(m\), the code computes \(n=m(m+1)\), extracts its squarefree part \(s\) and square part \(q\), computes the safe bound x_max, and stores the triple \((m,q,x_{\max})\) in a hash map keyed by \(s\).

The function find_fundamental_pell_limited uses continued fractions to find the fundamental Pell solution, and pell_solutions_up_to generates all further solutions by the standard Pell recurrence

$$ x_{t+1}=x_1x_t+s y_1y_t,\qquad y_{t+1}=x_1y_t+y_1x_t. $$

For each group with the same \(s\), the solver generates Pell solutions once up to the largest needed \(x\), then maps them back to candidate pivots with

$$ k=\frac{mx+qsy+m}{2}. $$

Finally it filters by oddness, lower bound, parity, and \(k\le L\), stores all hits in one vector, sorts them, removes duplicates, and sums them.

The checkpoint routine compares the fast method with brute force up to \(5000\) and also verifies that the sample pivots \(4,21,24,110\) appear below \(120\).

Complexity Analysis

The preprocessing range for \(m\) is only

$$ m\le O(\sqrt{L}), $$

and each such \(m\) is factorized once. The dominant work is then the Pell generation for each distinct squarefree class \(s\), together with the mapping and filtering of those solutions for every member of the group.

Memory usage is dominated by:

1. the grouping table keyed by squarefree part,

2. the temporary Pell solution list for one group,

3. the final vector of distinct pivots.

The critical practical idea is not a closed form for the entire problem, but the reduction from a three-parameter search \((k,m,n)\) to reusable Pell streams indexed by squarefree classes.

Footnotes and References

  1. Problem page: https://projecteuler.net/problem=261
  2. Sum of squares formula: Wikipedia - square pyramidal number
  3. Pell equations: Wikipedia - Pell's equation
  4. Continued fractions and Pell solvers: Wikipedia - continued fraction
  5. Squarefree integers: Wikipedia - square-free integer

Problem 261 source code

C++

#include <algorithm>
#include <cstdint>
#include <cmath>
#include <iostream>
#include <string>
#include <unordered_map>
#include <utility>
#include <vector>

namespace {

using i64 = long long;
using u64 = std::uint64_t;
using i128 = __int128_t;

struct Options {
    u64 limit = 10'000'000'000ULL;
    bool run_checkpoints = true;
};

struct MEntry {
    int m = 0;
    u64 q = 0;
    u64 x_max = 0;
};

struct PellSolution {
    u64 x = 0;
    u64 y = 0;
};

bool parse_u64_after_prefix(const std::string& arg, const std::string& prefix, u64& value) {
    if (arg.rfind(prefix, 0U) != 0U) {
        return false;
    }
    const std::string tail = arg.substr(prefix.size());
    if (tail.empty()) {
        return false;
    }
    u64 parsed = 0;
    for (char c : tail) {
        if (c < '0' || c > '9') {
            return false;
        }
        parsed = parsed * 10ULL + static_cast<u64>(c - '0');
    }
    value = parsed;
    return true;
}

bool parse_arguments(int argc, char** argv, Options& options) {
    for (int i = 1; i < argc; ++i) {
        const std::string arg(argv[i]);
        if (arg == "--skip-checkpoints") {
            options.run_checkpoints = false;
            continue;
        }
        if (parse_u64_after_prefix(arg, "--limit=", options.limit)) {
            continue;
        }
        std::cerr << "Unknown argument: " << arg << '\n';
        return false;
    }
    return options.limit >= 1;
}

std::vector<int> build_primes(const int n) {
    std::vector<int> spf(static_cast<std::size_t>(n + 1), 0);
    std::vector<int> primes;
    primes.reserve(n / 10);
    for (int i = 2; i <= n; ++i) {
        if (spf[static_cast<std::size_t>(i)] == 0) {
            spf[static_cast<std::size_t>(i)] = i;
            primes.push_back(i);
        }
        for (int p : primes) {
            const i64 v = static_cast<i64>(i) * p;
            if (v > n || p > spf[static_cast<std::size_t>(i)]) {
                break;
            }
            spf[static_cast<std::size_t>(v)] = p;
        }
    }
    return primes;
}

std::pair<u64, u64> squarefree_and_square_part(u64 n, const std::vector<int>& primes) {
    u64 s = 1;
    u64 q = 1;
    u64 x = n;

    for (int p : primes) {
        const u64 pp = static_cast<u64>(p);
        if (pp * pp > x) {
            break;
        }
        if (x % pp != 0) {
            continue;
        }
        int e = 0;
        while (x % pp == 0) {
            x /= pp;
            ++e;
        }
        if ((e & 1) != 0) {
            s *= pp;
        }
        for (int i = 0; i < e / 2; ++i) {
            q *= pp;
        }
    }

    if (x > 1) {
        s *= x;
    }

    return {s, q};
}

u64 floor_sqrt_u64(u64 x) {
    u64 r = static_cast<u64>(std::sqrt(static_cast<long double>(x)));
    while ((r + 1) * (r + 1) <= x) {
        ++r;
    }
    while (r * r > x) {
        --r;
    }
    return r;
}

bool find_fundamental_pell_limited(u64 s, u64 x_limit, u64& x1, u64& y1) {
    const u64 a0 = floor_sqrt_u64(s);
    if (a0 * a0 == s) {
        return false;
    }

    i128 m = 0;
    i128 d = 1;
    i128 a = static_cast<i128>(a0);

    i128 p_prev2 = 0;
    i128 p_prev1 = 1;
    i128 q_prev2 = 1;
    i128 q_prev1 = 0;

    for (int iter = 0; iter < 2'000'000; ++iter) {
        const i128 p = a * p_prev1 + p_prev2;
        const i128 q = a * q_prev1 + q_prev2;

        const i128 diff = p * p - static_cast<i128>(s) * q * q;
        if (diff == 1 && p > 1 && p <= static_cast<i128>(x_limit)) {
            x1 = static_cast<u64>(p);
            y1 = static_cast<u64>(q);
            return true;
        }

        if (p > static_cast<i128>(x_limit)) {
            return false;
        }

        p_prev2 = p_prev1;
        p_prev1 = p;
        q_prev2 = q_prev1;
        q_prev1 = q;

        m = d * a - m;
        d = (static_cast<i128>(s) - m * m) / d;
        a = (static_cast<i128>(a0) + m) / d;
    }

    return false;
}

std::vector<PellSolution> pell_solutions_up_to(u64 s, u64 x_limit) {
    u64 x1 = 0;
    u64 y1 = 0;
    if (!find_fundamental_pell_limited(s, x_limit, x1, y1)) {
        return {};
    }

    std::vector<PellSolution> out;
    u64 x = x1;
    u64 y = y1;

    while (x <= x_limit) {
        out.push_back({x, y});

        const i128 nx = static_cast<i128>(x1) * x + static_cast<i128>(s) * y1 * y;
        const i128 ny = static_cast<i128>(x1) * y + static_cast<i128>(y1) * x;
        if (nx > static_cast<i128>(x_limit)) {
            break;
        }
        x = static_cast<u64>(nx);
        y = static_cast<u64>(ny);
    }

    return out;
}

std::vector<u64> generate_square_pivots(const u64 limit) {
    const int m_max = static_cast<int>((std::sqrt(2.0L * static_cast<long double>(limit) + 1.0L) - 1.0L) / 2.0L);
    const std::vector<int> primes = build_primes(m_max + 2);

    std::unordered_map<u64, std::vector<MEntry>> grouped;
    grouped.reserve(static_cast<std::size_t>(m_max * 2));

    for (int m = 1; m <= m_max; ++m) {
        const u64 n = static_cast<u64>(m) * static_cast<u64>(m + 1);
        const auto [s, q] = squarefree_and_square_part(n, primes);
        const u64 x_max = (2ULL * limit - static_cast<u64>(m)) / static_cast<u64>(m);
        grouped[s].push_back({m, q, x_max});
    }

    std::vector<u64> pivots;
    pivots.reserve(1 << 20);

    for (auto& kv : grouped) {
        const u64 s = kv.first;
        std::vector<MEntry>& entries = kv.second;

        u64 x_max = 0;
        for (const MEntry& e : entries) {
            x_max = std::max(x_max, e.x_max);
        }

        const std::vector<PellSolution> sols = pell_solutions_up_to(s, x_max);
        if (sols.empty()) {
            continue;
        }

        for (const MEntry& e : entries) {
            const u64 min_d = static_cast<u64>(2 * e.m + 1);
            const i128 m128 = static_cast<i128>(e.m);
            const i128 qs128 = static_cast<i128>(e.q) * static_cast<i128>(s);

            for (const PellSolution& sol : sols) {
                if (sol.x < min_d) {
                    continue;
                }
                if ((sol.x & 1ULL) == 0ULL) {
                    continue;
                }

                const i128 numerator = m128 * static_cast<i128>(sol.x) +
                                       qs128 * static_cast<i128>(sol.y) + m128;
                if ((numerator & 1) != 0) {
                    continue;
                }

                const i128 k = numerator / 2;
                if (k > static_cast<i128>(limit)) {
                    continue;
                }
                if (k > 0) {
                    pivots.push_back(static_cast<u64>(k));
                }
            }
        }
    }

    std::sort(pivots.begin(), pivots.end());
    pivots.erase(std::unique(pivots.begin(), pivots.end()), pivots.end());
    return pivots;
}

u64 solve(const u64 limit) {
    const std::vector<u64> pivots = generate_square_pivots(limit);
    u64 sum = 0;
    for (u64 k : pivots) {
        sum += k;
    }
    return sum;
}

std::vector<u64> brute_pivots(const int limit) {
    std::vector<u64> pivots;
    for (int k = 1; k <= limit; ++k) {
        bool ok = false;
        for (int m = 1; m < k && !ok; ++m) {
            const i128 lhs = static_cast<i128>(m + 1) * k * (k - m);
            if (lhs % m != 0) {
                continue;
            }
            const i128 c = lhs / m;
            const i128 d = static_cast<i128>(m + 1) * (m + 1) + 4 * c;
            const i64 s = static_cast<i64>(std::sqrt(static_cast<long double>(d)));
            i64 root = s;
            while (static_cast<i128>(root + 1) * (root + 1) <= d) {
                ++root;
            }
            while (static_cast<i128>(root) * root > d) {
                --root;
            }
            if (static_cast<i128>(root) * root != d) {
                continue;
            }
            const i64 numer = -static_cast<i64>(m + 1) + root;
            if ((numer & 1) != 0) {
                continue;
            }
            const i64 n = numer / 2;
            if (n >= k) {
                ok = true;
            }
        }
        if (ok) {
            pivots.push_back(static_cast<u64>(k));
        }
    }
    return pivots;
}

bool run_checkpoints() {
    const auto fast_small = generate_square_pivots(5000);
    const auto brute_small = brute_pivots(5000);
    if (fast_small != brute_small) {
        std::cerr << "Checkpoint failed for limit=5000" << '\n';
        return false;
    }

    const auto sample = generate_square_pivots(120);
    const std::vector<u64> required = {4, 21, 24, 110};
    for (u64 x : required) {
        if (!std::binary_search(sample.begin(), sample.end(), x)) {
            std::cerr << "Checkpoint failed: missing sample pivot " << x << '\n';
            return false;
        }
    }

    return true;
}

}  // namespace

int main(int argc, char** argv) {
    Options options;
    if (!parse_arguments(argc, argv, options)) {
        return 1;
    }
    if (options.run_checkpoints && !run_checkpoints()) {
        return 2;
    }
    std::cout << solve(options.limit) << '\n';
    return 0;
}

Python

#!/usr/bin/env python3
import sys
import math

def build_primes(n):
    spf = [0] * (n + 1)
    primes = []
    for i in range(2, n + 1):
        if spf[i] == 0:
            spf[i] = i
            primes.append(i)
        for p in primes:
            v = i * p
            if v > n or p > spf[i]:
                break
            spf[v] = p
    return primes

def squarefree_and_square_part(n, primes):
    s = 1
    q = 1
    x = n
    for p in primes:
        pp = p
        if pp * pp > x:
            break
        if x % pp != 0:
            continue
        e = 0
        while x % pp == 0:
            x //= pp
            e += 1
        if e % 2 != 0:
            s *= pp
        for _ in range(e // 2):
            q *= pp
    if x > 1:
        s *= x
    return s, q

def floor_sqrt_u64(x):
    r = int(math.sqrt(float(x)))
    while (r + 1) * (r + 1) <= x:
        r += 1
    while r * r > x:
        r -= 1
    return r

def find_fundamental_pell_limited(s, x_limit):
    a0 = floor_sqrt_u64(s)
    if a0 * a0 == s:
        return None
    m = 0
    d = 1
    a = a0
    p_prev2 = 0
    p_prev1 = 1
    q_prev2 = 1
    q_prev1 = 0
    for _ in range(2_000_000):
        p = a * p_prev1 + p_prev2
        q = a * q_prev1 + q_prev2
        diff = p * p - s * q * q
        if diff == 1 and p > 1 and p <= x_limit:
            return p, q
        if p > x_limit:
            return None
        p_prev2 = p_prev1
        p_prev1 = p
        q_prev2 = q_prev1
        q_prev1 = q
        m = d * a - m
        d = (s - m * m) // d
        a = (a0 + m) // d
    return None

def pell_solutions_up_to(s, x_limit):
    res = find_fundamental_pell_limited(s, x_limit)
    if res is None:
        return []
    x1, y1 = res
    solutions = []
    x = x1
    y = y1
    while x <= x_limit:
        solutions.append((x, y))
        nx = x1 * x + s * y1 * y
        ny = x1 * y + y1 * x
        if nx > x_limit:
            break
        x = nx
        y = ny
    return solutions

def generate_square_pivots(limit):
    m_max = int((math.sqrt(2.0 * float(limit) + 1.0) - 1.0) / 2.0)
    primes = build_primes(m_max + 2)
    grouped = {}
    for m in range(1, m_max + 1):
        n = m * (m + 1)
        s, q = squarefree_and_square_part(n, primes)
        x_max = (2 * limit - m) // m
        if s not in grouped:
            grouped[s] = []
        grouped[s].append((m, q, x_max))
    pivots = []
    for s, entries in grouped.items():
        x_max = max(e[2] for e in entries)
        sols = pell_solutions_up_to(s, x_max)
        if not sols:
            continue
        for m, q, _ in entries:
            min_d = 2 * m + 1
            m128 = m
            qs128 = q * s
            for x, y in sols:
                if x < min_d:
                    continue
                if x % 2 == 0:
                    continue
                numerator = m128 * x + qs128 * y + m128
                if numerator % 2 != 0:
                    continue
                k = numerator // 2
                if k > limit:
                    continue
                if k > 0:
                    pivots.append(k)
    pivots.sort()
    return list(dict.fromkeys(pivots))

def solve(limit):
    pivots = generate_square_pivots(limit)
    return sum(pivots)

def brute_pivots(limit):
    pivots = []
    for k in range(1, limit + 1):
        ok = False
        for m in range(1, k):
            lhs = (m + 1) * k * (k - m)
            if lhs % m != 0:
                continue
            c = lhs // m
            d = (m + 1) * (m + 1) + 4 * c
            s = int(math.sqrt(float(d)))
            root = s
            while (root + 1) * (root + 1) <= d:
                root += 1
            while root * root > d:
                root -= 1
            if root * root != d:
                continue
            numer = -(m + 1) + root
            if numer % 2 != 0:
                continue
            n = numer // 2
            if n >= k:
                ok = True
                break
        if ok:
            pivots.append(k)
    return pivots

def run_checkpoints():
    fast_small = generate_square_pivots(5000)
    brute_small = brute_pivots(5000)
    if fast_small != brute_small:
        print("Checkpoint failed for limit=5000", file=sys.stderr)
        return False
    sample = generate_square_pivots(120)
    required = [4, 21, 24, 110]
    for x in required:
        if x not in sample:
            print(f"Checkpoint failed: missing sample pivot {x}", file=sys.stderr)
            return False
    return True

def parse_arguments(args):
    options = {
        "limit": 10_000_000_000,
        "run_checkpoints": True
    }
    for arg in args:
        if arg == "--skip-checkpoints":
            options["run_checkpoints"] = False
        elif arg.startswith("--limit="):
            try:
                options["limit"] = int(arg[8:])
            except ValueError:
                print(f"Unknown argument: {arg}", file=sys.stderr)
                return None
        else:
            print(f"Unknown argument: {arg}", file=sys.stderr)
            return None
    if options["limit"] < 1:
        return None
    return options

def main():
    args = sys.argv[1:]
    options = parse_arguments(args)
    if options is None:
        return 1
    if options["run_checkpoints"] and not run_checkpoints():
        return 2
    print(solve(options["limit"]))
    return 0

if __name__ == "__main__":
    sys.exit(main())

Java

import java.util.*;
import java.math.*;

class Euler261 {
    private static class Options {
        long limit = 10_000_000_000L;
        boolean runCheckpoints = true;
    }

    private static class MEntry {
        int m;
        long q;
        long xMax;
        
        MEntry(int m, long q, long xMax) {
            this.m = m;
            this.q = q;
            this.xMax = xMax;
        }
    }

    private static class PellSolution {
        long x;
        long y;
        
        PellSolution(long x, long y) {
            this.x = x;
            this.y = y;
        }
    }

    private static List<Integer> buildPrimes(int n) {
        int[] spf = new int[n + 1];
        List<Integer> primes = new ArrayList<>();
        for (int i = 2; i <= n; i++) {
            if (spf[i] == 0) {
                spf[i] = i;
                primes.add(i);
            }
            for (int p : primes) {
                long v = (long)i * p;
                if (v > n || p > spf[i]) break;
                spf[(int)v] = p;
            }
        }
        return primes;
    }

    private static long[] squarefreeAndSquarePart(long n, List<Integer> primes) {
        long s = 1;
        long q = 1;
        long x = n;
        for (int p : primes) {
            long pp = p;
            if (pp * pp > x) break;
            if (x % pp != 0) continue;
            int e = 0;
            while (x % pp == 0) {
                x /= pp;
                e++;
            }
            if ((e & 1) != 0) s *= pp;
            for (int i = 0; i < e/2; i++) q *= pp;
        }
        if (x > 1) s *= x;
        return new long[]{s, q};
    }

    private static long floorSqrtU64(long x) {
        long r = (long)Math.sqrt((double)x);
        while ((r + 1) * (r + 1) <= x) r++;
        while (r * r > x) r--;
        return r;
    }

    private static long[] findFundamentalPellLimited(long s, long xLimit) {
        long a0 = floorSqrtU64(s);
        if (a0 * a0 == s) return null;
        
        long m = 0;
        long d = 1;
        long a = a0;
        
        long pPrev2 = 0;
        long pPrev1 = 1;
        long qPrev2 = 1;
        long qPrev1 = 0;
        
        for (int iter = 0; iter < 2_000_000; iter++) {
            long p = a * pPrev1 + pPrev2;
            long q = a * qPrev1 + qPrev2;
            
            long diff = p * p - s * q * q;
            if (diff == 1 && p > 1 && p <= xLimit) {
                return new long[]{p, q};
            }
            
            if (p > xLimit) return null;
            
            pPrev2 = pPrev1;
            pPrev1 = p;
            qPrev2 = qPrev1;
            qPrev1 = q;
            
            m = d * a - m;
            d = (s - m * m) / d;
            a = (a0 + m) / d;
        }
        return null;
    }

    private static List<PellSolution> pellSolutionsUpTo(long s, long xLimit) {
        long[] fundamental = findFundamentalPellLimited(s, xLimit);
        if (fundamental == null) return new ArrayList<>();
        
        long x1 = fundamental[0];
        long y1 = fundamental[1];
        
        List<PellSolution> solutions = new ArrayList<>();
        long x = x1;
        long y = y1;
        
        while (x <= xLimit) {
            solutions.add(new PellSolution(x, y));
            BigInteger nx = BigInteger.valueOf(x1).multiply(BigInteger.valueOf(x))
                          .add(BigInteger.valueOf(s).multiply(BigInteger.valueOf(y1)).multiply(BigInteger.valueOf(y)));
            BigInteger ny = BigInteger.valueOf(x1).multiply(BigInteger.valueOf(y))
                          .add(BigInteger.valueOf(y1).multiply(BigInteger.valueOf(x)));
            if (nx.compareTo(BigInteger.valueOf(xLimit)) > 0) break;
            x = nx.longValue();
            y = ny.longValue();
        }
        return solutions;
    }

    private static List<Long> generateSquarePivots(long limit) {
        int mMax = (int)((Math.sqrt(2.0 * limit + 1.0) - 1.0) / 2.0);
        List<Integer> primes = buildPrimes(mMax + 2);
        Map<Long, List<MEntry>> grouped = new HashMap<>();
        
        for (int m = 1; m <= mMax; m++) {
            long n = (long)m * (m + 1);
            long[] sqParts = squarefreeAndSquarePart(n, primes);
            long s = sqParts[0];
            long q = sqParts[1];
            long xMax = (2 * limit - m) / m;
            grouped.computeIfAbsent(s, k -> new ArrayList<>()).add(new MEntry(m, q, xMax));
        }
        
        List<Long> pivots = new ArrayList<>();
        for (Map.Entry<Long, List<MEntry>> entry : grouped.entrySet()) {
            long s = entry.getKey();
            List<MEntry> entries = entry.getValue();
            
            long xMax = 0;
            for (MEntry e : entries) {
                xMax = Math.max(xMax, e.xMax);
            }
            
            List<PellSolution> sols = pellSolutionsUpTo(s, xMax);
            if (sols.isEmpty()) continue;
            
            for (MEntry e : entries) {
                long minD = 2L * e.m + 1;
                long m128 = e.m;
                long qs128 = e.q * s;
                
                for (PellSolution sol : sols) {
                    if (sol.x < minD) continue;
                    if ((sol.x & 1L) == 0L) continue;
                    
                    long numerator = m128 * sol.x + qs128 * sol.y + m128;
                    if ((numerator & 1L) != 0L) continue;
                    
                    long k = numerator / 2;
                    if (k > limit) continue;
                    if (k > 0) pivots.add(k);
                }
            }
        }
        
        Collections.sort(pivots);
        return new ArrayList<>(new LinkedHashSet<>(pivots));
    }

    private static long solve(long limit) {
        List<Long> pivots = generateSquarePivots(limit);
        long sum = 0;
        for (Long k : pivots) {
            sum += k;
        }
        return sum;
    }

    private static List<Long> brutePivots(int limit) {
        List<Long> pivots = new ArrayList<>();
        for (int k = 1; k <= limit; k++) {
            boolean ok = false;
            for (int m = 1; m < k && !ok; m++) {
                long lhs = (long)(m + 1) * k * (k - m);
                if (lhs % m != 0) continue;
                long c = lhs / m;
                long d = (long)(m + 1) * (m + 1) + 4 * c;
                long s = (long)Math.sqrt((double)d);
                long root = s;
                while ((root + 1) * (root + 1) <= d) root++;
                while (root * root > d) root--;
                if (root * root != d) continue;
                long numer = -(m + 1) + root;
                if ((numer & 1L) != 0) continue;
                long n = numer / 2;
                if (n >= k) ok = true;
            }
            if (ok) pivots.add((long)k);
        }
        return pivots;
    }

    private static boolean runCheckpoints() {
        List<Long> fastSmall = generateSquarePivots(5000);
        List<Long> bruteSmall = brutePivots(5000);
        if (!fastSmall.equals(bruteSmall)) {
            System.err.println("Checkpoint failed for limit=5000");
            return false;
        }
        
        List<Long> sample = generateSquarePivots(120);
        long[] required = {4, 21, 24, 110};
        for (long x : required) {
            if (!sample.contains(x)) {
                System.err.println("Checkpoint failed: missing sample pivot " + x);
                return false;
            }
        }
        return true;
    }

    private static Options parseArguments(String[] args) {
        Options options = new Options();
        for (String arg : args) {
            if (arg.equals("--skip-checkpoints")) {
                options.runCheckpoints = false;
            } else if (arg.startsWith("--limit=")) {
                try {
                    options.limit = Long.parseLong(arg.substring(8));
                } catch (NumberFormatException e) {
                    System.err.println("Unknown argument: " + arg);
                    return null;
                }
            } else {
                System.err.println("Unknown argument: " + arg);
                return null;
            }
        }
        if (options.limit < 1) return null;
        return options;
    }

    public static void main(String[] args) {
        Options options = parseArguments(args);
        if (options == null) {
            System.exit(1);
            return;
        }
        if (options.runCheckpoints && !runCheckpoints()) {
            System.exit(2);
            return;
        }
        System.out.println(solve(options.limit));
    }
}