Problem 660: Pandigital Triangles
View on Project EulerProject Euler Problem 660 Solution
EulerSolve provides an optimized solution for Project Euler Problem 660, Pandigital Triangles, with C++, Python, Java, and a step-by-step mathematical explanation.
Problem Summary For each base \(B\in[9,18]\), the digits \(0,1,\dots,B-1\) must be distributed exactly once across three positive integers \(a\), \(b\), and \(c\), with no leading zero. We seek all triples satisfying $$a^2+ab+b^2=c^2,\qquad c\lt a+b.$$ The equation is symmetric in \(a\) and \(b\), and the left-hand side is strictly larger than both \(a^2\) and \(b^2\), so every valid solution may be treated in ordered form \(a\lt b\lt c\). For each base we sum all valid values of \(c\), then add those sums over bases \(9\) through \(18\). Mathematical Approach The implementation constructs the numbers from left to right. A search state contains a current prefix for each of the three numbers and the set of digits already used. The key mathematics is what lets the search discard entire subtrees long before all digits are assigned. Step 1: Organize the search by balanced digit lengths Write $$B=3k+r,\qquad r\in\{0,1,2\}.$$ The solver arranges the final digit lengths as $$ (\ell_a,\ell_b,\ell_c)= \begin{cases} (k,k,k), & r=0,\\ (k,k,k+1), & r=1,\\ (k,k+1,k+1), & r=2. \end{cases} $$ This is the balanced way to spend all \(B\) digits, and it matches the magnitude restriction \(b\lt c\lt a+b\): \(c\) must be slightly larger than \(b\), but it cannot be dramatically larger....
Detailed mathematical approach
Problem Summary
For each base \(B\in[9,18]\), the digits \(0,1,\dots,B-1\) must be distributed exactly once across three positive integers \(a\), \(b\), and \(c\), with no leading zero. We seek all triples satisfying
$$a^2+ab+b^2=c^2,\qquad c\lt a+b.$$
The equation is symmetric in \(a\) and \(b\), and the left-hand side is strictly larger than both \(a^2\) and \(b^2\), so every valid solution may be treated in ordered form \(a\lt b\lt c\). For each base we sum all valid values of \(c\), then add those sums over bases \(9\) through \(18\).
Mathematical Approach
The implementation constructs the numbers from left to right. A search state contains a current prefix for each of the three numbers and the set of digits already used. The key mathematics is what lets the search discard entire subtrees long before all digits are assigned.
Step 1: Organize the search by balanced digit lengths
Write
$$B=3k+r,\qquad r\in\{0,1,2\}.$$
The solver arranges the final digit lengths as
$$ (\ell_a,\ell_b,\ell_c)= \begin{cases} (k,k,k), & r=0,\\ (k,k,k+1), & r=1,\\ (k,k+1,k+1), & r=2. \end{cases} $$
This is the balanced way to spend all \(B\) digits, and it matches the magnitude restriction \(b\lt c\lt a+b\): \(c\) must be slightly larger than \(b\), but it cannot be dramatically larger. The search therefore starts with one leading digit in each number when \(r=0\), with one extra leading digit already placed in \(c\) when \(r=1\), and with one extra leading digit already placed in both \(b\) and \(c\) when \(r=2\).
Step 2: Extend prefixes while preserving order and digit uniqueness
If the current prefixes are still called \(a\), \(b\), and \(c\), then appending one new base-\(B\) digit to each gives
$$a'=aB+d_a,\qquad b'=bB+d_b,\qquad c'=cB+d_c,$$
where \(d_a\), \(d_b\), and \(d_c\) are distinct unused digits. Only the leading digits are forced to be nonzero; later appended digits may be zero.
The ordering established at the start is preserved automatically. Indeed, if \(a\lt b\), then for any digits \(d_a,d_b\in\{0,\dots,B-1\}\),
$$aB+d_a\le aB+(B-1) \lt (a+1)B\le bB\le bB+d_b.$$
So once a prefix state is generated in sorted order, every deeper state stays sorted as well. This prevents duplicate counting caused by the symmetry between \(a\) and \(b\).
Step 3: Use interval bounds to test whether a prefix can still succeed
Suppose there are \(r\) more digits to append to each number, and set
$$p=B^r.$$
Then the final completed values must lie in the intervals
$$A\in[ap,\ ap+p-1],\qquad B\in[bp,\ bp+p-1],\qquad C\in[cp,\ cp+p-1].$$
Because the quadratic form \(x^2+xy+y^2\) is increasing for positive \(x\) and \(y\), the smallest and largest possible left-hand sides over this state are
$$L_{\min}=(ap)^2+(bp)^2+(ap)(bp),$$
$$L_{\max}=(ap+p-1)^2+(bp+p-1)^2+(ap+p-1)(bp+p-1).$$
Likewise, the right-hand side must stay inside
$$R_{\min}=(cp)^2,\qquad R_{\max}=(cp+p-1)^2.$$
A prefix can be extended only if two necessary conditions hold:
$$cp \lt (ap+p-1)+(bp+p-1),$$
$$L_{\max}\ge R_{\min},\qquad L_{\min}\le R_{\max}.$$
If either the triangle inequality range or the quadratic-form range fails, then no completion of that prefix can possibly work, so the entire branch is pruned immediately.
Step 4: Finish exactly when only two or three digits per number remain
Once the search has reached the point where only two digits remain for each number, there are exactly six unused digits left. At that point the implementation simply tests all \(6!=720\) ordered assignments of those digits to the three 2-digit tails and checks the full equation exactly.
When three digits remain for each number, the implementation uses a stronger filter. Write the final numbers as
$$A=aB^3+a_t,\qquad B=bB^3+b_t,\qquad C=cB^3+c_t,$$
where \(a_t\), \(b_t\), and \(c_t\) are 3-digit tails built from the remaining digits. If
$$A^2+AB+B^2=C^2,$$
then reducing modulo \(B^3\) yields the necessary congruence
$$a_t^2+a_tb_t+b_t^2\equiv c_t^2\pmod{B^3}.$$
So the solver precomputes all ordered 3-digit tails with distinct digits, groups candidate \(c_t\) values by their square residue modulo \(B^3\), and only combines tails whose digit masks are disjoint and whose residues satisfy the congruence. The full identity and \(c\lt a+b\) are then checked on the completed integers.
Worked Example: pruning a base-9 prefix
Take \(B=9\). Then \(9=3\cdot3\), so the final lengths are \((3,3,3)\). Suppose a search node has one-digit prefixes
$$a=1,\qquad b=2,\qquad c=5,$$
with two more digits still to append to each number. Then \(p=9^2=81\), so
$$A\in[81,161],\qquad B\in[162,242],\qquad C\in[405,485].$$
Now
$$\max A+\max B=161+242=403,$$
but
$$\min C=405.$$
Therefore every completion of this prefix would satisfy \(C\ge A+B\), which contradicts \(c\lt a+b\). The whole subtree below \((1,2,5)\) is discarded without trying any of the remaining six digits.
How the Code Works
The C++ and Java implementations follow the same search strategy. For each base they precompute powers of the base, generate the allowed initial prefix states for the relevant length pattern, and repeatedly append one digit to each number while keeping only the states that survive the interval test above.
If the search reaches the 2-digit finishing case, the implementation checks all \(720\) tail assignments directly. If it reaches the 3-digit finishing case, it precomputes all distinct 3-digit tails, buckets them by residue modulo \(B^3\), joins only mask-compatible residue matches, and then verifies the complete identity with exact integer arithmetic rather than floating point. The C++ implementation also parallelizes the independent bases \(9\) through \(18\) and includes an independent brute-force verification for base \(9\). The Python implementation is intentionally thin: it builds or reuses the compiled solver, runs it, and returns the parsed numeric result.
Complexity Analysis
A naive search would be essentially factorial in \(B\), because it would have to examine permutations of the \(B\) digits together with ways to split them among three numbers. The implemented method still has exponential worst-case behavior, but its practical cost is driven by the number of prefix states that survive pruning rather than by raw \(B!\).
For the finishing phase, the 2-digit case costs \(O(720)\) exact completions per live state. In the 3-digit case there are
$$P(B,3)=B(B-1)(B-2)$$
ordered tails to precompute, and the residue-based preprocessing is quadratic in that tail count up to the number of residue-compatible, mask-disjoint pairs and triples that survive. Memory usage is proportional to the live prefix states plus the stored tail tables. In the C++ version, distributing bases across worker threads improves wall-clock time but does not change the underlying search complexity.
Footnotes and References
- Problem page: https://projecteuler.net/problem=660
- Pandigital numbers: Wikipedia — Pandigital number
- Positional notation: Wikipedia — Positional notation
- Branch and bound: Wikipedia — Branch and bound
- Eisenstein integers and the related norm form: Wikipedia — Eisenstein integer
Problem 660 source code
C++
#include <algorithm>
#include <array>
#include <atomic>
#include <cassert>
#include <cstdint>
#include <iostream>
#include <pthread.h>
#include <unordered_set>
#include <unistd.h>
#include <vector>
namespace {
using u32 = std::uint32_t;
using u64 = std::uint64_t;
using u128 = unsigned __int128;
struct State {
u64 a;
u64 b;
u64 c;
u32 mask;
};
struct Tail {
int value;
u32 mask;
};
struct TailTriple {
int a;
int b;
int c;
};
bool feasible(u64 a, u64 b, u64 c, u64 p) {
const u64 min_a = a * p;
const u64 min_b = b * p;
const u64 min_c = c * p;
const u64 max_a = min_a + p - 1;
const u64 max_b = min_b + p - 1;
const u64 max_c = min_c + p - 1;
if (min_c >= max_a + max_b) return false;
const u128 max_left = static_cast<u128>(max_a) * max_a +
static_cast<u128>(max_b) * max_b +
static_cast<u128>(max_a) * max_b;
const u128 min_left = static_cast<u128>(min_a) * min_a +
static_cast<u128>(min_b) * min_b +
static_cast<u128>(min_a) * min_b;
const u128 min_right = static_cast<u128>(min_c) * min_c;
const u128 max_right = static_cast<u128>(max_c) * max_c;
if (max_left < min_right) return false;
if (min_left > max_right) return false;
return true;
}
std::vector<int> digits_from_mask(u32 mask, int base) {
std::vector<int> digits;
digits.reserve(static_cast<std::size_t>(base));
for (int d = 0; d < base; ++d) {
if ((mask & (1u << d)) == 0) digits.push_back(d);
}
return digits;
}
std::vector<State> initial_states(int base, u64 p) {
std::vector<State> states;
const int r = base % 3;
if (r == 0) {
for (int a = 1; a < base; ++a) {
for (int b = a + 1; b < base; ++b) {
for (int c = b + 1; c < base; ++c) {
u32 mask = (1u << a) | (1u << b) | (1u << c);
if (feasible(a, b, c, p)) {
states.push_back({static_cast<u64>(a), static_cast<u64>(b), static_cast<u64>(c), mask});
}
}
}
}
return states;
}
if (r == 1) {
for (int a = 1; a < base; ++a) {
for (int b = a + 1; b < base; ++b) {
u32 mask_ab = (1u << a) | (1u << b);
for (int c1 = 1; c1 < base; ++c1) {
if (mask_ab & (1u << c1)) continue;
for (int c0 = 0; c0 < base; ++c0) {
if (c0 == c1) continue;
if (mask_ab & (1u << c0)) continue;
const int c = c1 * base + c0;
const u32 mask = mask_ab | (1u << c1) | (1u << c0);
if (feasible(a, b, c, p)) {
states.push_back({static_cast<u64>(a), static_cast<u64>(b), static_cast<u64>(c), mask});
}
}
}
}
}
return states;
}
for (int a = 1; a < base; ++a) {
const u32 mask_a = (1u << a);
for (int b1 = 1; b1 < base; ++b1) {
if (mask_a & (1u << b1)) continue;
for (int b0 = 0; b0 < base; ++b0) {
if (b0 == b1) continue;
u32 mask_ab = mask_a | (1u << b1) | (1u << b0);
if (__builtin_popcount(mask_ab) != 3) continue;
const int b = b1 * base + b0;
for (int c1 = 1; c1 < base; ++c1) {
if (mask_ab & (1u << c1)) continue;
for (int c0 = 0; c0 < base; ++c0) {
if (c0 == c1) continue;
if (mask_ab & (1u << c0)) continue;
const int c = c1 * base + c0;
if (b >= c) continue;
const u32 mask = mask_ab | (1u << c1) | (1u << c0);
if (feasible(a, b, c, p)) {
states.push_back({static_cast<u64>(a), static_cast<u64>(b), static_cast<u64>(c), mask});
}
}
}
}
}
}
return states;
}
std::vector<State> extend_states(const std::vector<State>& states, int base, u64 p) {
std::vector<State> next;
for (const auto& s : states) {
const auto digits = digits_from_mask(s.mask, base);
const int m = static_cast<int>(digits.size());
for (int ia = 0; ia < m; ++ia) {
const int da = digits[ia];
for (int ib = 0; ib < m; ++ib) {
if (ib == ia) continue;
const int db = digits[ib];
for (int ic = 0; ic < m; ++ic) {
if (ic == ia || ic == ib) continue;
const int dc = digits[ic];
const u32 mask = s.mask | (1u << da) | (1u << db) | (1u << dc);
const u64 na = s.a * base + static_cast<u64>(da);
const u64 nb = s.b * base + static_cast<u64>(db);
const u64 nc = s.c * base + static_cast<u64>(dc);
if (feasible(na, nb, nc, p)) {
next.push_back({na, nb, nc, mask});
}
}
}
}
}
return next;
}
u64 finalize_two_digits(const std::vector<State>& states, int base) {
const u64 pow2 = static_cast<u64>(base) * static_cast<u64>(base);
u64 sum = 0;
for (const auto& s : states) {
auto digits = digits_from_mask(s.mask, base);
if (digits.size() != 6) continue;
std::sort(digits.begin(), digits.end());
do {
const u64 a = s.a * pow2 + static_cast<u64>(digits[0] * base + digits[1]);
const u64 b = s.b * pow2 + static_cast<u64>(digits[2] * base + digits[3]);
const u64 c = s.c * pow2 + static_cast<u64>(digits[4] * base + digits[5]);
const u128 left = static_cast<u128>(a) * a + static_cast<u128>(b) * b + static_cast<u128>(a) * b;
const u128 right = static_cast<u128>(c) * c;
if (left == right && c < a + b) sum += c;
} while (std::next_permutation(digits.begin(), digits.end()));
}
return sum;
}
u64 finalize_three_digits(const std::vector<State>& states, int base) {
const int mod = base * base * base;
const u64 pow3 = static_cast<u64>(base) * base * base;
std::vector<Tail> tails;
tails.reserve(static_cast<std::size_t>(base * (base - 1) * (base - 2)));
for (int d2 = 0; d2 < base; ++d2) {
for (int d1 = 0; d1 < base; ++d1) {
if (d1 == d2) continue;
for (int d0 = 0; d0 < base; ++d0) {
if (d0 == d1 || d0 == d2) continue;
const int value = d2 * base * base + d1 * base + d0;
const u32 mask = (1u << d2) | (1u << d1) | (1u << d0);
tails.push_back({value, mask});
}
}
}
std::vector<std::vector<int>> c_by_res(static_cast<std::size_t>(mod));
for (std::size_t i = 0; i < tails.size(); ++i) {
const u64 v = static_cast<u64>(tails[i].value);
const int res = static_cast<int>((v * v) % static_cast<u64>(mod));
c_by_res[static_cast<std::size_t>(res)].push_back(static_cast<int>(i));
}
std::vector<std::vector<TailTriple>> by_mask(1u << base);
for (const auto& a : tails) {
for (const auto& b : tails) {
if (a.mask & b.mask) continue;
const u64 av = static_cast<u64>(a.value);
const u64 bv = static_cast<u64>(b.value);
const int res = static_cast<int>((av * av + bv * bv + av * bv) % static_cast<u64>(mod));
const u32 ab_mask = a.mask | b.mask;
for (int idx : c_by_res[static_cast<std::size_t>(res)]) {
const auto& c = tails[static_cast<std::size_t>(idx)];
if (ab_mask & c.mask) continue;
const u32 mask = ab_mask | c.mask;
by_mask[mask].push_back({a.value, b.value, c.value});
}
}
}
const u32 full_mask = (1u << base) - 1u;
u64 sum = 0;
for (const auto& s : states) {
const u32 remaining = full_mask ^ s.mask;
for (const auto& t : by_mask[remaining]) {
const u64 a = s.a * pow3 + static_cast<u64>(t.a);
const u64 b = s.b * pow3 + static_cast<u64>(t.b);
const u64 c = s.c * pow3 + static_cast<u64>(t.c);
const u128 left = static_cast<u128>(a) * a + static_cast<u128>(b) * b + static_cast<u128>(a) * b;
const u128 right = static_cast<u128>(c) * c;
if (left == right && c < a + b) sum += c;
}
}
return sum;
}
u64 solve_base(int base) {
const int k = base / 3;
const int remaining = k - 1;
std::vector<u64> pow_base(static_cast<std::size_t>(k + 1), 1);
for (int i = 1; i <= k; ++i) pow_base[static_cast<std::size_t>(i)] = pow_base[static_cast<std::size_t>(i - 1)] * base;
auto states = initial_states(base, pow_base[static_cast<std::size_t>(remaining)]);
if (remaining > 3) {
for (int rem = remaining - 1; rem >= 3; --rem) {
states = extend_states(states, base, pow_base[static_cast<std::size_t>(rem)]);
}
}
if (remaining == 2) return finalize_two_digits(states, base);
return finalize_three_digits(states, base);
}
u64 brute_base9() {
std::array<int, 9> digits{};
for (int i = 0; i < 9; ++i) digits[static_cast<std::size_t>(i)] = i;
struct TripleKey {
u64 a;
u64 b;
u64 c;
};
struct Hash {
std::size_t operator()(const TripleKey& t) const noexcept {
std::size_t h = static_cast<std::size_t>(t.a);
h ^= static_cast<std::size_t>(t.b) + 0x9e3779b97f4a7c15ULL + (h << 6) + (h >> 2);
h ^= static_cast<std::size_t>(t.c) + 0x9e3779b97f4a7c15ULL + (h << 6) + (h >> 2);
return h;
}
};
struct Eq {
bool operator()(const TripleKey& x, const TripleKey& y) const noexcept {
return x.a == y.a && x.b == y.b && x.c == y.c;
}
};
std::unordered_set<TripleKey, Hash, Eq> seen;
u64 sum = 0;
do {
for (int i = 1; i <= 7; ++i) {
if (digits[0] == 0) continue;
for (int j = i + 1; j <= 8; ++j) {
if (digits[static_cast<std::size_t>(i)] == 0 || digits[static_cast<std::size_t>(j)] == 0) continue;
u64 a = 0;
for (int x = 0; x < i; ++x) a = a * 9 + static_cast<u64>(digits[static_cast<std::size_t>(x)]);
u64 b = 0;
for (int x = i; x < j; ++x) b = b * 9 + static_cast<u64>(digits[static_cast<std::size_t>(x)]);
u64 c = 0;
for (int x = j; x < 9; ++x) c = c * 9 + static_cast<u64>(digits[static_cast<std::size_t>(x)]);
std::array<u64, 3> s{a, b, c};
std::sort(s.begin(), s.end());
a = s[0];
b = s[1];
c = s[2];
const u128 left = static_cast<u128>(a) * a + static_cast<u128>(b) * b + static_cast<u128>(a) * b;
const u128 right = static_cast<u128>(c) * c;
if (left != right || c >= a + b) continue;
const TripleKey key{a, b, c};
if (seen.insert(key).second) sum += c;
}
}
} while (std::next_permutation(digits.begin(), digits.end()));
return sum;
}
struct WorkerData {
std::atomic<int>* next_base;
std::vector<u64>* results;
};
void* worker_fn(void* arg) {
auto* data = static_cast<WorkerData*>(arg);
while (true) {
const int base = data->next_base->fetch_add(1);
if (base > 18) break;
(*data->results)[static_cast<std::size_t>(base)] = solve_base(base);
}
return nullptr;
}
} // namespace
int main() {
std::vector<u64> results(19, 0);
std::atomic<int> next_base{9};
long threads = sysconf(_SC_NPROCESSORS_ONLN);
if (threads < 1) threads = 1;
if (threads > 10) threads = 10;
std::vector<pthread_t> pool(static_cast<std::size_t>(threads));
WorkerData data{&next_base, &results};
for (long i = 0; i < threads; ++i) {
pthread_create(&pool[static_cast<std::size_t>(i)], nullptr, worker_fn, &data);
}
for (long i = 0; i < threads; ++i) {
pthread_join(pool[static_cast<std::size_t>(i)], nullptr);
}
const u64 brute9 = brute_base9();
assert(results[9] == brute9);
u64 sum = 0;
for (int base = 9; base <= 18; ++base) sum += results[static_cast<std::size_t>(base)];
std::cout << sum << '\n';
return 0;
}
Python
from __future__ import annotations
import re
import shutil
import subprocess
from pathlib import Path
ANSWER_RE = re.compile(r"answer\s*:\s*(.+)$", re.IGNORECASE)
EQUAL_RE = re.compile(r"=\s*(.+)$")
def parse_output(stdout: str) -> str:
lines = [line.strip() for line in stdout.splitlines() if line.strip()]
if not lines:
return ""
answers = []
equals = []
for line in lines:
m1 = ANSWER_RE.search(line)
if m1:
answers.append(m1.group(1).strip())
m2 = EQUAL_RE.search(line)
if m2:
equals.append(m2.group(1).strip())
if answers:
return answers[-1]
if equals:
return equals[-1]
return lines[-1]
def should_skip_cpp_checkpoints(src: Path) -> bool:
try:
text = src.read_text(encoding="utf-8", errors="ignore")
except OSError:
return False
return "--skip-checkpoints" in text
def run_cpp(binary: Path, src: Path, root: Path) -> str:
cmd = [str(binary)]
if should_skip_cpp_checkpoints(src):
cmd.append("--skip-checkpoints")
try:
return subprocess.check_output(cmd, text=True, cwd=root)
except subprocess.CalledProcessError:
return subprocess.check_output(cmd, text=True, cwd=src.parent)
def solve() -> str:
problem_id = __file__.split("Euler")[-1].split(".")[0]
root = Path(__file__).resolve().parent.parent
src = root / "solutionsCpp" / f"Euler{problem_id}.cpp"
binary = root / "solutionsCpp" / f".euler{problem_id}_py_bridge"
if not binary.exists() or src.stat().st_mtime > binary.stat().st_mtime:
compiler = shutil.which("clang++") or shutil.which("g++")
if not compiler:
raise RuntimeError("No C++ compiler found (clang++/g++).")
subprocess.check_call([compiler, "-std=c++17", "-O2", str(src), "-o", str(binary)])
output = run_cpp(binary=binary, src=src, root=root)
parsed = parse_output(output)
if not parsed:
raise RuntimeError(f"Euler{problem_id} bridge produced empty output.")
return parsed
if __name__ == "__main__":
print(solve())
Java
import java.util.*;
public class Euler660 {
public static String solve() {
long total = 0;
for (int base = 9; base <= 18; base++) {
int k = base / 3, rem = k - 1;
long[] powBase = new long[k + 1];
powBase[0] = 1;
for (int i = 1; i <= k; i++)
powBase[i] = powBase[i - 1] * base;
List<long[]> states = new ArrayList<>();
int r = base % 3;
long p = powBase[rem];
if (r == 0) {
for (int a = 1; a < base; a++)
for (int b = a + 1; b < base; b++)
for (int c = b + 1; c < base; c++) {
int mask = (1 << a) | (1 << b) | (1 << c);
if (feasible(a, b, c, p, base))
states.add(new long[] { a, b, c, mask });
}
} else if (r == 1) {
for (int a = 1; a < base; a++)
for (int b = a + 1; b < base; b++) {
int mab = (1 << a) | (1 << b);
for (int c1 = 1; c1 < base; c1++)
if ((mab & (1 << c1)) == 0)
for (int c0 = 0; c0 < base; c0++) {
if (c0 == c1 || (mab & (1 << c0)) != 0)
continue;
long cv = c1 * base + c0;
int mask = mab | (1 << c1) | (1 << c0);
if (feasible(a, b, cv, p, base))
states.add(new long[] { a, b, cv, mask });
}
}
} else {
for (int a = 1; a < base; a++) {
int ma = 1 << a;
for (int b1 = 1; b1 < base; b1++) {
if ((ma & (1 << b1)) != 0)
continue;
for (int b0 = 0; b0 < base; b0++) {
if (b0 == b1)
continue;
int mab = ma | (1 << b1) | (1 << b0);
if (Integer.bitCount(mab) != 3)
continue;
long bv = b1 * base + b0;
for (int c1 = 1; c1 < base; c1++) {
if ((mab & (1 << c1)) != 0)
continue;
for (int c0 = 0; c0 < base; c0++) {
if (c0 == c1 || (mab & (1 << c0)) != 0)
continue;
long cv = c1 * base + c0;
if (bv >= cv)
continue;
int mask = mab | (1 << c1) | (1 << c0);
if (feasible(a, bv, cv, p, base))
states.add(new long[] { a, bv, cv, mask });
}
}
}
}
}
}
for (int step = rem - 1; step > (rem > 3 ? 2 : rem - 1); step--) {
List<long[]> nxt = new ArrayList<>();
long pp = powBase[step];
for (long[] st : states) {
long a = st[0], b = st[1], c = st[2];
int mask = (int) st[3];
for (int da = 0; da < base; da++) {
if ((mask & (1 << da)) != 0)
continue;
for (int db = 0; db < base; db++) {
if (db == da || (mask & (1 << db)) != 0)
continue;
for (int dc = 0; dc < base; dc++) {
if (dc == da || dc == db || (mask & (1 << dc)) != 0)
continue;
int nm = mask | (1 << da) | (1 << db) | (1 << dc);
long na = a * base + da, nb = b * base + db, nc = c * base + dc;
if (feasible(na, nb, nc, pp, base))
nxt.add(new long[] { na, nb, nc, nm });
}
}
}
}
states = nxt;
}
if (rem == 2) {
long p2 = (long) base * base;
for (long[] st : states) {
long a = st[0], b = st[1], c = st[2];
int mask = (int) st[3];
List<Integer> digits = new ArrayList<>();
for (int d = 0; d < base; d++)
if ((mask & (1 << d)) == 0)
digits.add(d);
if (digits.size() != 6)
continue;
int[] dd = digits.stream().mapToInt(Integer::intValue).toArray();
for (int[] perm : permutations6(dd)) {
long fa = a * p2 + perm[0] * base + perm[1], fb = b * p2 + perm[2] * base + perm[3],
fc = c * p2 + perm[4] * base + perm[5];
if (fa * fa + fb * fb + fa * fb == fc * fc && fc < fa + fb)
total += fc;
}
}
} else {
long p3 = (long) base * base * base;
int full = (1 << base) - 1;
Map<Integer, List<int[]>> byMask = new HashMap<>();
for (int d2 = 0; d2 < base; d2++)
for (int d1 = 0; d1 < base; d1++) {
if (d1 == d2)
continue;
for (int d0 = 0; d0 < base; d0++) {
if (d0 == d1 || d0 == d2)
continue;
int v = d2 * base * base + d1 * base + d0, m = (1 << d2) | (1 << d1) | (1 << d0);
// Pre-group by residue? Just store all tails
byMask.computeIfAbsent(m, x -> new ArrayList<>()).add(new int[] { v, m });
}
}
// Simpler brute approach for remaining
List<int[]> tails = new ArrayList<>();
for (int d2 = 0; d2 < base; d2++)
for (int d1 = 0; d1 < base; d1++) {
if (d1 == d2)
continue;
for (int d0 = 0; d0 < base; d0++) {
if (d0 == d1 || d0 == d2)
continue;
tails.add(
new int[] { d2 * base * base + d1 * base + d0, (1 << d2) | (1 << d1) | (1 << d0) });
}
}
long mod3 = p3;
Map<Long, List<Integer>> cByRes = new HashMap<>();
for (int i = 0; i < tails.size(); i++)
cByRes.computeIfAbsent((long) tails.get(i)[0] * tails.get(i)[0] % mod3, x -> new ArrayList<>())
.add(i);
Map<Integer, List<long[]>> byMask2 = new HashMap<>();
for (int ai = 0; ai < tails.size(); ai++) {
int av = tails.get(ai)[0], am = tails.get(ai)[1];
for (int bi = 0; bi < tails.size(); bi++) {
int bv = tails.get(bi)[0], bm = tails.get(bi)[1];
if ((am & bm) != 0)
continue;
long res = ((long) av * av + (long) bv * bv + (long) av * bv) % mod3;
int abm = am | bm;
List<Integer> cis = cByRes.get(res);
if (cis == null)
continue;
for (int ci : cis) {
int cv = tails.get(ci)[0], cm = tails.get(ci)[1];
if ((abm & cm) != 0)
continue;
int mk = abm | cm;
byMask2.computeIfAbsent(mk, x -> new ArrayList<>()).add(new long[] { av, bv, cv });
}
}
}
for (long[] st : states) {
long a = st[0], b = st[1], c = st[2];
int mask = (int) st[3];
int remaining = full ^ mask;
List<long[]> entries = byMask2.get(remaining);
if (entries == null)
continue;
for (long[] e : entries) {
long fa = a * p3 + e[0], fb = b * p3 + e[1], fc = c * p3 + e[2];
if (fa * fa + fb * fb + fa * fb == fc * fc && fc < fa + fb)
total += fc;
}
}
}
}
return String.valueOf(total);
}
static boolean feasible(long a, long b, long c, long p, int base) {
long mnA = a * p, mxA = a * p + p - 1, mnB = b * p, mxB = b * p + p - 1, mnC = c * p, mxC = c * p + p - 1;
if (mnC >= mxA + mxB)
return false;
long ml = mxA * mxA + mxB * mxB + mxA * mxB, mnl = mnA * mnA + mnB * mnB + mnA * mnB;
return ml >= mnC * mnC && mnl <= mxC * mxC;
}
static int[][] permutations6(int[] arr) {
List<int[]> result = new ArrayList<>();
perm6(arr, 0, result);
return result.toArray(new int[0][]);
}
static void perm6(int[] arr, int idx, List<int[]> result) {
if (idx == arr.length) {
result.add(arr.clone());
return;
}
for (int i = idx; i < arr.length; i++) {
int t = arr[i];
arr[i] = arr[idx];
arr[idx] = t;
perm6(arr, idx + 1, result);
arr[idx] = arr[i];
arr[i] = t;
}
}
public static void main(String[] args) {
System.out.println(solve());
}
}