Problem 991: Fruit Salad
View on Project EulerProject Euler Problem 991 Solution
EulerSolve provides an optimized solution for Project Euler Problem 991, Fruit Salad, with C++, Python, Java, and a step-by-step mathematical explanation.
Problem Summary Let the three fruit quantities be positive integers \(a,b,c\). In the published statement, the second and third fractions share the same denominator \(a+c\), so the condition is $$\frac{a}{b+c}+\frac{b}{a+c}+\frac{c}{a+c}=4,$$ which immediately simplifies to $$\frac{a}{b+c}+\frac{b+c}{a+c}=4.$$ For every solution with perimeter \(p=a+b+c\le 10^7\), we must add \(p\) to the final total. A direct search over all triples would be hopeless, so the real task is to classify the integer solutions in a structured way. Mathematical Approach The implementations work by identifying primitive solutions, turning them into explicit parameter families, and then summing all multiples of each primitive perimeter at once. From the Rational Identity to a Discriminant Multiplying by \((a+c)(b+c)\) removes the denominators: $$a(a+c)+(b+c)^2=4(a+c)(b+c).$$ After expansion this becomes the quadratic Diophantine equation $$a^2-4ab-3ac+b^2-2bc-3c^2=0.$$ Now treat it as a quadratic in \(b\): $$b^2-(4a+2c)b+(a^2-3ac-3c^2)=0.$$ Its discriminant is $$\Delta=(4a+2c)^2-4(a^2-3ac-3c^2)=4(a+c)(3a+4c).$$ Therefore any integer solution must satisfy $$b=2a+c\pm \sqrt{(a+c)(3a+4c)},$$ so \((a+c)(3a+4c)\) has to be a perfect square. Primitive Reduction and the Key Coprimality Fact After clearing denominators, the equation is homogeneous of degree 2....
Detailed mathematical approach
Problem Summary
Let the three fruit quantities be positive integers \(a,b,c\). In the published statement, the second and third fractions share the same denominator \(a+c\), so the condition is
$$\frac{a}{b+c}+\frac{b}{a+c}+\frac{c}{a+c}=4,$$
which immediately simplifies to
$$\frac{a}{b+c}+\frac{b+c}{a+c}=4.$$
For every solution with perimeter \(p=a+b+c\le 10^7\), we must add \(p\) to the final total. A direct search over all triples would be hopeless, so the real task is to classify the integer solutions in a structured way.
Mathematical Approach
The implementations work by identifying primitive solutions, turning them into explicit parameter families, and then summing all multiples of each primitive perimeter at once.
From the Rational Identity to a Discriminant
Multiplying by \((a+c)(b+c)\) removes the denominators:
$$a(a+c)+(b+c)^2=4(a+c)(b+c).$$
After expansion this becomes the quadratic Diophantine equation
$$a^2-4ab-3ac+b^2-2bc-3c^2=0.$$
Now treat it as a quadratic in \(b\):
$$b^2-(4a+2c)b+(a^2-3ac-3c^2)=0.$$
Its discriminant is
$$\Delta=(4a+2c)^2-4(a^2-3ac-3c^2)=4(a+c)(3a+4c).$$
Therefore any integer solution must satisfy
$$b=2a+c\pm \sqrt{(a+c)(3a+4c)},$$
so \((a+c)(3a+4c)\) has to be a perfect square.
Primitive Reduction and the Key Coprimality Fact
After clearing denominators, the equation is homogeneous of degree 2. If \((a,b,c)\) is a solution, then so is \(k(a,b,c)\) for every positive integer \(k\). That means every solution is a multiple of a primitive one with \(\gcd(a,b,c)=1\).
For a primitive triple we only need one coprimality statement: \(\gcd(a,c)=1\). Indeed, if some \(d\) divides both \(a\) and \(c\), then in
$$a^2-4ab-3ac+b^2-2bc-3c^2=0$$
every term except \(b^2\) is divisible by \(d\), hence \(d\mid b^2\), so \(d\mid b\). That would force \(d\mid \gcd(a,b,c)\), contradicting primitiveness.
Now compute
$$\gcd(a+c,3a+4c)=\gcd(a+c,\;3a+4c-3(a+c))=\gcd(a+c,c)=\gcd(a,c)=1.$$
So the two factors in \((a+c)(3a+4c)\) are coprime.
Replacing One Square by Two Squares
We already know that \((a+c)(3a+4c)\) is a square, and the previous step shows the two factors are coprime. By unique factorization, each factor must therefore be a square on its own:
$$a+c=u^2,\qquad 3a+4c=v^2,\qquad \gcd(u,v)=1.$$
Solving this linear system gives
$$a=4u^2-v^2,\qquad c=v^2-3u^2.$$
Substituting back into the quadratic formula for \(b\) yields
$$b=5u^2-v^2\pm uv.$$
Thus every primitive solution has the form
$$\bigl(a,b,c\bigr)=\bigl(4u^2-v^2,\;5u^2-v^2\pm uv,\;v^2-3u^2\bigr),\qquad \gcd(u,v)=1.$$
Positivity and Perimeter Factorization
Since \(a\gt 0\) and \(c\gt 0\), we obtain
$$v^2\lt 4u^2,\qquad v^2\gt 3u^2,$$
so the valid region is the narrow strip
$$\sqrt{3}\,u \lt v \lt 2u.$$
The primitive perimeter is
$$p_\pm=a+b+c=6u^2-v^2\pm uv.$$
These expressions factor neatly:
$$p_-=(2u-v)(3u+v),\qquad p_+=(2u+v)(3u-v).$$
Those two factorizations are the bridge from the theoretical parametrization to the compact formulas used in the implementations.
Family I: \(u\) and \(v\) of Opposite Parity
Because \(\gcd(u,v)=1\), the pair cannot be both even. The first implemented family covers the case where \(u\) and \(v\) have opposite parity.
For the perimeter \(p_-\), set
$$m=u,\qquad n=2u-v,$$
so \(v=2m-n\). Then
$$a=4mn-n^2,\qquad b=5mn-m^2-n^2,\qquad c=m^2-4mn+n^2,$$
and
$$p_-=n(5m-n).$$
For the perimeter \(p_+\), set
$$m=2u+v,\qquad n=u,$$
so \(v=m-2n\). This gives
$$a=4mn-m^2,\qquad b=5mn-m^2-n^2,\qquad c=m^2-4mn+n^2,$$
and
$$p_+=m(5n-m).$$
In both branches we get \(\gcd(m,n)=1\) and \(m-n\) odd. The positivity conditions on \(b\) and \(c\) imply
$$2+\sqrt{3}\lt \frac{m}{n}\lt \frac{5+\sqrt{21}}{2}\lt 5,$$
which explains why the implementations can safely scan only up to \(m\le 5n\) in this family.
Family II: \(u\) and \(v\) Both Odd
The only remaining primitive possibility is that both \(u\) and \(v\) are odd.
For the perimeter \(p_-\), define
$$m=\frac{3u-v}{2},\qquad n=\frac{v-u}{2},$$
so equivalently
$$u=m+n,\qquad v=m+3n.$$
Then
$$a=3m^2+2mn-5n^2,\qquad b=3m^2-7n^2,\qquad c=2(3n^2-m^2),$$
and
$$p_-=4m^2+2mn-6n^2.$$
For the perimeter \(p_+\), define
$$m=\frac{3u+v}{2},\qquad n=\frac{u+v}{2},$$
so equivalently
$$u=m-n,\qquad v=3n-m.$$
This yields
$$a=3m^2-2mn-5n^2,\qquad b=3m^2-7n^2,\qquad c=2(3n^2-m^2),$$
and
$$p_+=4m^2-2mn-6n^2.$$
Again \(\gcd(m,n)=1\) and \(m-n\) is odd. Here the common positivity conditions on \(b\) and \(c\) give
$$\sqrt{\frac{7}{3}}\lt \frac{m}{n}\lt \sqrt{3}\lt 2,$$
so the scan only needs \(m\le 2n\) in this second family.
Worked Example
Take the smallest primitive pair that works: \(u=4\), \(v=7\). Then
$$a=4\cdot 4^2-7^2=15,\qquad c=7^2-3\cdot 4^2=1.$$
For \(b\) we obtain
$$b=5\cdot 4^2-7^2\pm 4\cdot 7=31\pm 28\in\{3,59\}.$$
So the two primitive solutions are \((15,3,1)\) with perimeter \(19\), and \((15,59,1)\) with perimeter \(75\). If the search bound were \(N=50\), only the first triple and its double \((30,6,2)\) would contribute, producing
$$19+38=57.$$
This is exactly the small checkpoint reproduced by the implementations.
How the Code Works
Bounding the Search
For any primitive solution, \(u^2=a+c\lt a+b+c=p\le N\). So the fundamental parameters are only \(O(\sqrt{N})\), and the linear substitutions above keep the implementation parameters \(m,n\) in the same range.
Enumerating the Two Families
The C++, Python, and Java implementations loop over positive \(n\), then scan \(m\) over the two ranges implied by the derivation: up to \(5n\) for the opposite-parity family and up to \(2n\) for the odd-odd family. They enforce \(\gcd(m,n)=1\) and opposite parity, compute the corresponding formulas for \(a\), \(b\), \(c\), and discard any candidate with a nonpositive component.
Adding Every Multiple in One Step
If a primitive triple has perimeter \(p\), then every multiple \(kp\) is also valid as long as \(kp\le N\). With
$$q=\left\lfloor\frac{N}{p}\right\rfloor,$$
the total contribution of that primitive triple is
$$p+2p+\cdots+qp=p\frac{q(q+1)}{2}.$$
That closed form is why the implementations never need to generate each multiple explicitly. They also include a small brute-force verifier for low bounds, and the C++ implementation can optionally split the outer parameter range across several threads.
Complexity Analysis
Let \(N\) be the perimeter limit. The outer parameter only grows like \(\sqrt{N}\). Across all \(n\), the first family examines \(O\!\left(\sum n\right)=O(N)\) candidate pairs \((m,n)\), and the second family does the same. Each candidate uses only fixed-width integer arithmetic plus one gcd test, so in ordinary RAM-model terms the running time is \(O(N)\).
Memory usage is \(O(1)\) aside from a few counters and, in the multithreaded C++ version, a small array of partial sums. For \(N=10^7\) this leaves only a few million arithmetic checks, which is easily manageable.
Footnotes and References
- Problem page: https://projecteuler.net/problem=991
- Diophantine equation: Wikipedia - Diophantine equation
- Coprime integers: Wikipedia - Coprime integers
- Fundamental theorem of arithmetic: Wikipedia - Fundamental theorem of arithmetic
- Quadratic equation and discriminant: Wikipedia - Quadratic equation
- Arithmetic progression: Wikipedia - Arithmetic progression
Problem 991 source code
C++
#include <algorithm>
#include <cstdint>
#include <iostream>
#include <numeric>
#include <string>
#include <thread>
#include <vector>
namespace {
using i64 = std::int64_t;
using u64 = std::uint64_t;
using u128 = unsigned __int128;
struct Options {
u64 limit = 10'000'000ULL;
bool run_checkpoints = true;
bool allow_multithreading = true;
unsigned requested_threads = 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 = 0ULL;
for (const char c : tail) {
if (c < '0' || c > '9') {
return false;
}
const u64 digit = static_cast<u64>(c - '0');
if (parsed > (std::numeric_limits<u64>::max() - digit) / 10ULL) {
return false;
}
parsed = parsed * 10ULL + digit;
}
value = parsed;
return true;
}
bool parse_unsigned_after_prefix(const std::string& arg, const std::string& prefix, unsigned& value) {
u64 parsed = 0ULL;
if (!parse_u64_after_prefix(arg, prefix, parsed) || parsed == 0ULL ||
parsed > static_cast<u64>(std::numeric_limits<unsigned>::max())) {
return false;
}
value = static_cast<unsigned>(parsed);
return true;
}
bool parse_arguments(const 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 (arg == "--no-mt") {
options.allow_multithreading = false;
continue;
}
if (parse_u64_after_prefix(arg, "--limit=", options.limit)) {
continue;
}
if (parse_unsigned_after_prefix(arg, "--threads=", options.requested_threads)) {
continue;
}
std::cerr << "Unknown argument: " << arg << '\n';
return false;
}
return true;
}
u64 isqrt_u64(const u64 n) {
u64 low = 0ULL;
u64 high = std::min<u64>(n, 1ULL << 32);
while (low < high) {
const u64 mid = low + (high - low + 1ULL) / 2ULL;
if (mid <= n / mid) {
low = mid;
} else {
high = mid - 1ULL;
}
}
return low;
}
u128 arithmetic_series_sum(const u64 count) {
return static_cast<u128>(count) * static_cast<u128>(count + 1ULL) / 2U;
}
u128 contribution(const u64 limit, const u64 base_sum) {
const u64 copies = limit / base_sum;
return static_cast<u128>(base_sum) * arithmetic_series_sum(copies);
}
u128 solve_range(const u64 limit, const u64 n_begin, const u64 n_end) {
u128 total = 0;
for (u64 n = n_begin; n < n_end; ++n) {
const i64 nn = static_cast<i64>(n);
// Case 1: A = k * 2mn, B = k * (m^2 - n^2), C = k * (m^2 + n^2).
for (u64 m = n + 1ULL; m <= 5ULL * n; ++m) {
if (((m - n) & 1ULL) == 0ULL || std::gcd(m, n) != 1ULL) {
continue;
}
const i64 mm = static_cast<i64>(m);
const i64 b = 5LL * mm * nn - mm * mm - nn * nn;
const i64 c = mm * mm - 4LL * mm * nn + nn * nn;
if (b <= 0LL || c <= 0LL) {
continue;
}
const u64 sum_plus = static_cast<u64>(nn * (5LL * mm - nn));
if (sum_plus <= limit) {
total += contribution(limit, sum_plus);
}
const i64 a_minus = 4LL * mm * nn - mm * mm;
if (a_minus > 0LL) {
const u64 sum_minus = static_cast<u64>(mm * (5LL * nn - mm));
if (sum_minus <= limit) {
total += contribution(limit, sum_minus);
}
}
}
// Case 2: A = 2k * (m^2 - n^2), B = 4k mn, C = 2k * (m^2 + n^2).
for (u64 m = n + 1ULL; m <= 2ULL * n; ++m) {
if (((m - n) & 1ULL) == 0ULL || std::gcd(m, n) != 1ULL) {
continue;
}
const i64 mm = static_cast<i64>(m);
const i64 b = 3LL * mm * mm - 7LL * nn * nn;
const i64 c = 2LL * (3LL * nn * nn - mm * mm);
if (b <= 0LL || c <= 0LL) {
continue;
}
const i64 sum_plus = 4LL * mm * mm + 2LL * mm * nn - 6LL * nn * nn;
if (sum_plus > 0LL && static_cast<u64>(sum_plus) <= limit) {
total += contribution(limit, static_cast<u64>(sum_plus));
}
const i64 a_minus = 3LL * mm * mm - 2LL * mm * nn - 5LL * nn * nn;
if (a_minus > 0LL) {
const i64 sum_minus = 4LL * mm * mm - 2LL * mm * nn - 6LL * nn * nn;
if (sum_minus > 0LL && static_cast<u64>(sum_minus) <= limit) {
total += contribution(limit, static_cast<u64>(sum_minus));
}
}
}
}
return total;
}
unsigned resolve_thread_count(const Options& options, const u64 max_n) {
if (!options.allow_multithreading || max_n <= 1ULL) {
return 1U;
}
if (options.requested_threads != 0U) {
return std::min<unsigned>(options.requested_threads, static_cast<unsigned>(max_n));
}
const unsigned detected = std::thread::hardware_concurrency();
const unsigned fallback = detected == 0U ? 1U : detected;
return std::min<unsigned>(fallback, static_cast<unsigned>(max_n));
}
u128 solve(const u64 limit, const unsigned thread_count) {
if (limit == 0ULL) {
return 0;
}
const u64 max_n = isqrt_u64(limit) + 2ULL;
if (thread_count <= 1U || max_n <= 1ULL) {
return solve_range(limit, 1ULL, max_n + 1ULL);
}
std::vector<u128> partials(thread_count, 0);
std::vector<std::thread> workers;
workers.reserve(thread_count);
const u64 block = (max_n + static_cast<u64>(thread_count) - 1ULL) / static_cast<u64>(thread_count);
for (unsigned tid = 0; tid < thread_count; ++tid) {
const u64 begin = 1ULL + static_cast<u64>(tid) * block;
const u64 end = std::min<u64>(max_n + 1ULL, begin + block);
workers.emplace_back([&, tid, begin, end]() {
if (begin < end) {
partials[tid] = solve_range(limit, begin, end);
}
});
}
for (auto& worker : workers) {
worker.join();
}
u128 total = 0;
for (const u128 part : partials) {
total += part;
}
return total;
}
u128 brute_force(const u64 limit) {
u128 total = 0;
for (u64 a = 1ULL; a <= limit; ++a) {
for (u64 b = 1ULL; a + b <= limit; ++b) {
for (u64 c = 1ULL; a + b + c <= limit; ++c) {
const u128 lhs = static_cast<u128>(a) * static_cast<u128>(a + c) +
static_cast<u128>(b + c) * static_cast<u128>(b + c);
const u128 rhs =
4U * static_cast<u128>(b + c) * static_cast<u128>(a + c);
if (lhs == rhs) {
total += static_cast<u128>(a + b + c);
}
}
}
}
return total;
}
std::string to_string_u128(u128 value) {
if (value == 0) {
return "0";
}
std::string digits;
while (value != 0) {
const unsigned digit = static_cast<unsigned>(value % 10U);
digits.push_back(static_cast<char>('0' + digit));
value /= 10U;
}
std::reverse(digits.begin(), digits.end());
return digits;
}
bool run_checkpoints() {
if (solve(50ULL, 1U) != brute_force(50ULL)) {
std::cerr << "Checkpoint failed for limit=50" << '\n';
return false;
}
if (solve(200ULL, 1U) != brute_force(200ULL)) {
std::cerr << "Checkpoint failed for limit=200" << '\n';
return false;
}
if (solve(10'000'000ULL, 1U) != static_cast<u128>(23'871'972'654'940ULL)) {
std::cerr << "Checkpoint failed for limit=10000000" << '\n';
return false;
}
return true;
}
} // namespace
int main(int argc, char** argv) {
Options options;
if (!parse_arguments(argc, argv, options)) {
return 1;
}
const u64 max_n = isqrt_u64(options.limit) + 2ULL;
const unsigned thread_count = resolve_thread_count(options, max_n);
if (options.run_checkpoints && !run_checkpoints()) {
return 2;
}
std::cout << to_string_u128(solve(options.limit, thread_count)) << '\n';
return 0;
}
Python
import math
import sys
DEFAULT_LIMIT = 10_000_000
def parse_arguments(argv):
limit = DEFAULT_LIMIT
run_checkpoints = True
for arg in argv[1:]:
if arg == "--skip-checkpoints":
run_checkpoints = False
continue
if arg.startswith("--limit="):
tail = arg[len("--limit="):]
if not tail.isdigit():
raise ValueError(f"Invalid limit: {arg}")
limit = int(tail)
continue
raise ValueError(f"Unknown argument: {arg}")
return limit, run_checkpoints
def arithmetic_series_sum(count):
return count * (count + 1) // 2
def contribution(limit, base_sum):
copies = limit // base_sum
return base_sum * arithmetic_series_sum(copies)
def solve(limit=DEFAULT_LIMIT):
if limit <= 0:
return 0
gcd = math.gcd
total = 0
max_n = math.isqrt(limit) + 2
for n in range(1, max_n + 1):
for m in range(n + 1, 5 * n + 1, 2):
if gcd(m, n) != 1:
continue
b = 5 * m * n - m * m - n * n
c = m * m - 4 * m * n + n * n
if b <= 0 or c <= 0:
continue
sum_plus = n * (5 * m - n)
if sum_plus <= limit:
total += contribution(limit, sum_plus)
a_minus = 4 * m * n - m * m
if a_minus > 0:
sum_minus = m * (5 * n - m)
if sum_minus <= limit:
total += contribution(limit, sum_minus)
for m in range(n + 1, 2 * n + 1, 2):
if gcd(m, n) != 1:
continue
b = 3 * m * m - 7 * n * n
c = 2 * (3 * n * n - m * m)
if b <= 0 or c <= 0:
continue
sum_plus = 4 * m * m + 2 * m * n - 6 * n * n
if sum_plus > 0 and sum_plus <= limit:
total += contribution(limit, sum_plus)
a_minus = 3 * m * m - 2 * m * n - 5 * n * n
if a_minus > 0:
sum_minus = 4 * m * m - 2 * m * n - 6 * n * n
if sum_minus > 0 and sum_minus <= limit:
total += contribution(limit, sum_minus)
return total
def brute_force(limit):
total = 0
for a in range(1, limit + 1):
for b in range(1, limit - a + 1):
max_c = limit - a - b
for c in range(1, max_c + 1):
if a * (a + c) + (b + c) * (b + c) == 4 * (b + c) * (a + c):
total += a + b + c
return total
def run_checkpoints():
if solve(50) != brute_force(50):
raise AssertionError("Checkpoint failed for limit=50")
if solve(200) != brute_force(200):
raise AssertionError("Checkpoint failed for limit=200")
def main(argv=None):
argv = sys.argv if argv is None else argv
try:
limit, should_run_checkpoints = parse_arguments(argv)
if should_run_checkpoints:
run_checkpoints()
print(solve(limit))
except ValueError as exc:
print(exc, file=sys.stderr)
return 1
except AssertionError as exc:
print(exc, file=sys.stderr)
return 2
return 0
if __name__ == "__main__":
raise SystemExit(main())
Java
import java.math.BigInteger;
public class Euler991 {
private static final long DEFAULT_LIMIT = 10_000_000L;
private static final class Options {
long limit = DEFAULT_LIMIT;
boolean runCheckpoints = true;
}
private static boolean parseLongAfterPrefix(String arg, String prefix, Options options) {
if (!arg.startsWith(prefix)) {
return false;
}
String tail = arg.substring(prefix.length());
if (tail.isEmpty()) {
throw new IllegalArgumentException("Invalid limit: " + arg);
}
long value = 0L;
for (int i = 0; i < tail.length(); ++i) {
char ch = tail.charAt(i);
if (ch < '0' || ch > '9') {
throw new IllegalArgumentException("Invalid limit: " + arg);
}
int digit = ch - '0';
if (value > (Long.MAX_VALUE - digit) / 10L) {
throw new IllegalArgumentException("Limit overflow: " + arg);
}
value = value * 10L + digit;
}
options.limit = value;
return true;
}
private static Options parseArguments(String[] args) {
Options options = new Options();
for (String arg : args) {
if ("--skip-checkpoints".equals(arg)) {
options.runCheckpoints = false;
continue;
}
if (parseLongAfterPrefix(arg, "--limit=", options)) {
continue;
}
throw new IllegalArgumentException("Unknown argument: " + arg);
}
return options;
}
private static long isqrt(long n) {
long low = 0L;
long high = Math.min(n, 1L << 32);
while (low < high) {
long mid = low + (high - low + 1L) / 2L;
if (mid <= n / mid) {
low = mid;
} else {
high = mid - 1L;
}
}
return low;
}
private static BigInteger arithmeticSeriesSum(long count) {
return BigInteger.valueOf(count).multiply(BigInteger.valueOf(count + 1L)).divide(BigInteger.TWO);
}
private static BigInteger contribution(long limit, long baseSum) {
long copies = limit / baseSum;
return BigInteger.valueOf(baseSum).multiply(arithmeticSeriesSum(copies));
}
private static long gcd(long a, long b) {
while (b != 0L) {
long t = a % b;
a = b;
b = t;
}
return a;
}
private static BigInteger solve(long limit) {
if (limit <= 0L) {
return BigInteger.ZERO;
}
BigInteger total = BigInteger.ZERO;
long maxN = isqrt(limit) + 2L;
for (long n = 1L; n <= maxN; ++n) {
for (long m = n + 1L; m <= 5L * n; m += 2L) {
if (gcd(m, n) != 1L) {
continue;
}
long b = 5L * m * n - m * m - n * n;
long c = m * m - 4L * m * n + n * n;
if (b <= 0L || c <= 0L) {
continue;
}
long sumPlus = n * (5L * m - n);
if (sumPlus <= limit) {
total = total.add(contribution(limit, sumPlus));
}
long aMinus = 4L * m * n - m * m;
if (aMinus > 0L) {
long sumMinus = m * (5L * n - m);
if (sumMinus <= limit) {
total = total.add(contribution(limit, sumMinus));
}
}
}
for (long m = n + 1L; m <= 2L * n; m += 2L) {
if (gcd(m, n) != 1L) {
continue;
}
long b = 3L * m * m - 7L * n * n;
long c = 2L * (3L * n * n - m * m);
if (b <= 0L || c <= 0L) {
continue;
}
long sumPlus = 4L * m * m + 2L * m * n - 6L * n * n;
if (sumPlus > 0L && sumPlus <= limit) {
total = total.add(contribution(limit, sumPlus));
}
long aMinus = 3L * m * m - 2L * m * n - 5L * n * n;
if (aMinus > 0L) {
long sumMinus = 4L * m * m - 2L * m * n - 6L * n * n;
if (sumMinus > 0L && sumMinus <= limit) {
total = total.add(contribution(limit, sumMinus));
}
}
}
}
return total;
}
private static BigInteger bruteForce(long limit) {
BigInteger total = BigInteger.ZERO;
for (long a = 1L; a <= limit; ++a) {
for (long b = 1L; a + b <= limit; ++b) {
for (long c = 1L; a + b + c <= limit; ++c) {
long lhs = a * (a + c) + (b + c) * (b + c);
long rhs = 4L * (b + c) * (a + c);
if (lhs == rhs) {
total = total.add(BigInteger.valueOf(a + b + c));
}
}
}
}
return total;
}
private static void runCheckpoints() {
if (!solve(50L).equals(bruteForce(50L))) {
throw new IllegalStateException("Checkpoint failed for limit=50");
}
if (!solve(200L).equals(bruteForce(200L))) {
throw new IllegalStateException("Checkpoint failed for limit=200");
}
}
public static void main(String[] args) {
try {
Options options = parseArguments(args);
if (options.runCheckpoints) {
runCheckpoints();
}
System.out.println(solve(options.limit));
} catch (IllegalArgumentException ex) {
System.err.println(ex.getMessage());
System.exit(1);
} catch (IllegalStateException ex) {
System.err.println(ex.getMessage());
System.exit(2);
}
}
}