diff --git a/quartz/src/commonMain/kotlin/com/vitorpamplona/quartz/utils/secp256k1/Fe4.kt b/quartz/src/commonMain/kotlin/com/vitorpamplona/quartz/utils/secp256k1/Fe4.kt new file mode 100644 index 000000000..0e4714bd4 --- /dev/null +++ b/quartz/src/commonMain/kotlin/com/vitorpamplona/quartz/utils/secp256k1/Fe4.kt @@ -0,0 +1,93 @@ +/* + * Copyright (c) 2025 Vitor Pamplona + * + * Permission is hereby granted, free of charge, to any person obtaining a copy of + * this software and associated documentation files (the "Software"), to deal in + * the Software without restriction, including without limitation the rights to use, + * copy, modify, merge, publish, distribute, sublicense, and/or sell copies of the + * Software, and to permit persons to whom the Software is furnished to do so, + * subject to the following conditions: + * + * The above copyright notice and this permission notice shall be included in all + * copies or substantial portions of the Software. + * + * THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR + * IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, FITNESS + * FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE AUTHORS OR + * COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER LIABILITY, WHETHER IN + * AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, OUT OF OR IN CONNECTION + * WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE SOFTWARE. + */ +package com.vitorpamplona.quartz.utils.secp256k1 + +import kotlin.jvm.JvmField + +// ===================================================================================== +// STRUCT-BASED FIELD ELEMENT: 4 named Long fields instead of LongArray(4). +// +// MOTIVATION: +// LongArray(4) requires a bounds check on every laload/lastore instruction in JVM +// bytecode. Even though HotSpot C2 can often eliminate these for constant indices +// when the array size is known, the situation is worse on: +// - Android ART: Limited inlining depth means fewer bounds checks eliminated. +// - Kotlin/Native LLVM AOT: Parameter arrays are harder to analyze statically. +// +// With named fields (l0..l3), access compiles to direct getfield/putfield bytecode +// on JVM, and direct struct field access on Native — zero bounds checks on any +// platform, ever. +// +// MEMORY: +// LongArray(4): 12-byte header + 4-byte length + 32 bytes data = 48 bytes +// Fe4 object: 16-byte header + 32 bytes data = 48 bytes (identical) +// +// BYTECODE COMPARISON (per field access): +// LongArray: aload + iconst + laload (3 insns + implicit bounds check) +// Fe4: aload + getfield (2 insns, no check) +// +// ===================================================================================== + +/** + * Mutable 256-bit field element using 4 named Long fields in little-endian order. + * l0 = least significant 64 bits, l3 = most significant 64 bits. + * + * @JvmField eliminates virtual getter/setter generation — direct field access. + */ +internal class Fe4( + @JvmField var l0: Long = 0L, + @JvmField var l1: Long = 0L, + @JvmField var l2: Long = 0L, + @JvmField var l3: Long = 0L, +) { + fun isZero(): Boolean = (l0 or l1 or l2 or l3) == 0L + + fun copyFrom(other: Fe4) { + l0 = other.l0 + l1 = other.l1 + l2 = other.l2 + l3 = other.l3 + } + + fun setZero() { + l0 = 0L + l1 = 0L + l2 = 0L + l3 = 0L + } +} + +/** + * Mutable 512-bit wide product buffer using 8 named Long fields. + * Used for intermediate results of 256×256-bit multiplication before reduction. + * + * Replaces LongArray(8) scratch buffers — eliminates 8 bounds checks per access. + */ +internal class Wide8( + @JvmField var l0: Long = 0L, + @JvmField var l1: Long = 0L, + @JvmField var l2: Long = 0L, + @JvmField var l3: Long = 0L, + @JvmField var l4: Long = 0L, + @JvmField var l5: Long = 0L, + @JvmField var l6: Long = 0L, + @JvmField var l7: Long = 0L, +) diff --git a/quartz/src/commonMain/kotlin/com/vitorpamplona/quartz/utils/secp256k1/Fe4Ops.kt b/quartz/src/commonMain/kotlin/com/vitorpamplona/quartz/utils/secp256k1/Fe4Ops.kt new file mode 100644 index 000000000..9f662f886 --- /dev/null +++ b/quartz/src/commonMain/kotlin/com/vitorpamplona/quartz/utils/secp256k1/Fe4Ops.kt @@ -0,0 +1,599 @@ +/* + * Copyright (c) 2025 Vitor Pamplona + * + * Permission is hereby granted, free of charge, to any person obtaining a copy of + * this software and associated documentation files (the "Software"), to deal in + * the Software without restriction, including without limitation the rights to use, + * copy, modify, merge, publish, distribute, sublicense, and/or sell copies of the + * Software, and to permit persons to whom the Software is furnished to do so, + * subject to the following conditions: + * + * The above copyright notice and this permission notice shall be included in all + * copies or substantial portions of the Software. + * + * THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR + * IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, FITNESS + * FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE AUTHORS OR + * COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER LIABILITY, WHETHER IN + * AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, OUT OF OR IN CONNECTION + * WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE SOFTWARE. + */ +package com.vitorpamplona.quartz.utils.secp256k1 + +// ===================================================================================== +// FIELD OPERATIONS ON Fe4 (struct-based, zero bounds checks) +// +// Mirror implementations of U256 and FieldP operations using Fe4/Wide8 instead of +// LongArray(4)/LongArray(8). The arithmetic is identical; only the access pattern +// differs (field access vs array indexing). +// +// These are used for benchmarking to measure the impact of eliminating array bounds +// checks on various platforms (JVM HotSpot, Android ART, Kotlin/Native LLVM). +// ===================================================================================== + +/** + * 256-bit unsigned arithmetic on Fe4 (struct-based, no bounds checks). + */ +internal object Fe4U256 { + /** out = a + b. Returns carry (0 or 1). */ + fun addTo( + out: Fe4, + a: Fe4, + b: Fe4, + ): Int { + var s1: Long + var s2: Long + var c1: Long + var c2: Long + + // Limb 0 + s1 = a.l0 + b.l0 + c1 = if (uLtInline(s1, a.l0)) 1L else 0L + out.l0 = s1 + var carry = c1 + + // Limb 1 + s1 = a.l1 + b.l1 + c1 = if (uLtInline(s1, a.l1)) 1L else 0L + s2 = s1 + carry + c2 = if (uLtInline(s2, s1)) 1L else 0L + out.l1 = s2 + carry = c1 + c2 + + // Limb 2 + s1 = a.l2 + b.l2 + c1 = if (uLtInline(s1, a.l2)) 1L else 0L + s2 = s1 + carry + c2 = if (uLtInline(s2, s1)) 1L else 0L + out.l2 = s2 + carry = c1 + c2 + + // Limb 3 + s1 = a.l3 + b.l3 + c1 = if (uLtInline(s1, a.l3)) 1L else 0L + s2 = s1 + carry + c2 = if (uLtInline(s2, s1)) 1L else 0L + out.l3 = s2 + carry = c1 + c2 + + return carry.toInt() + } + + /** out = a - b. Returns borrow (0 or 1). */ + fun subTo( + out: Fe4, + a: Fe4, + b: Fe4, + ): Int { + var d1: Long + var d2: Long + var c1: Long + var c2: Long + + // Limb 0 + d1 = a.l0 - b.l0 + c1 = if (uLtInline(a.l0, b.l0)) 1L else 0L + out.l0 = d1 + var borrow = c1 + + // Limb 1 + d1 = a.l1 - b.l1 + c1 = if (uLtInline(a.l1, b.l1)) 1L else 0L + d2 = d1 - borrow + c2 = if (uLtInline(d1, borrow)) 1L else 0L + out.l1 = d2 + borrow = c1 + c2 + + // Limb 2 + d1 = a.l2 - b.l2 + c1 = if (uLtInline(a.l2, b.l2)) 1L else 0L + d2 = d1 - borrow + c2 = if (uLtInline(d1, borrow)) 1L else 0L + out.l2 = d2 + borrow = c1 + c2 + + // Limb 3 + d1 = a.l3 - b.l3 + c1 = if (uLtInline(a.l3, b.l3)) 1L else 0L + d2 = d1 - borrow + c2 = if (uLtInline(d1, borrow)) 1L else 0L + out.l3 = d2 + borrow = c1 + c2 + + return borrow.toInt() + } + + /** 4×4 schoolbook multiplication: w = a × b (512-bit result). */ + @Suppress("LongMethod") + fun mulWide( + w: Wide8, + a: Fe4, + b: Fe4, + ) { + val a0 = a.l0 + val a1 = a.l1 + val a2 = a.l2 + val a3 = a.l3 + val b0 = b.l0 + val b1 = b.l1 + val b2 = b.l2 + val b3 = b.l3 + var lo: Long + var hi: Long + var prev: Long + var s: Long + var c1: Long + var c2: Long + var carry: Long + + // Row 0: a0 × [b0,b1,b2,b3] + lo = a0 * b0 + w.l0 = lo + carry = unsignedMultiplyHigh(a0, b0) + + lo = a0 * b1 + s = lo + carry + c1 = if (uLtInline(s, lo)) 1L else 0L + w.l1 = s + carry = unsignedMultiplyHigh(a0, b1) + c1 + + lo = a0 * b2 + s = lo + carry + c1 = if (uLtInline(s, lo)) 1L else 0L + w.l2 = s + carry = unsignedMultiplyHigh(a0, b2) + c1 + + lo = a0 * b3 + s = lo + carry + c1 = if (uLtInline(s, lo)) 1L else 0L + w.l3 = s + w.l4 = unsignedMultiplyHigh(a0, b3) + c1 + + // Row 1: a1 × [b0,b1,b2,b3] + lo = a1 * b0 + hi = unsignedMultiplyHigh(a1, b0) + prev = w.l1 + s = prev + lo + c1 = if (uLtInline(s, prev)) 1L else 0L + w.l1 = s + carry = hi + c1 + + lo = a1 * b1 + hi = unsignedMultiplyHigh(a1, b1) + prev = w.l2 + s = prev + lo + c1 = if (uLtInline(s, prev)) 1L else 0L + s += carry + c2 = if (uLtInline(s, carry)) 1L else 0L + w.l2 = s + carry = hi + c1 + c2 + + lo = a1 * b2 + hi = unsignedMultiplyHigh(a1, b2) + prev = w.l3 + s = prev + lo + c1 = if (uLtInline(s, prev)) 1L else 0L + s += carry + c2 = if (uLtInline(s, carry)) 1L else 0L + w.l3 = s + carry = hi + c1 + c2 + + lo = a1 * b3 + hi = unsignedMultiplyHigh(a1, b3) + prev = w.l4 + s = prev + lo + c1 = if (uLtInline(s, prev)) 1L else 0L + s += carry + c2 = if (uLtInline(s, carry)) 1L else 0L + w.l4 = s + w.l5 = hi + c1 + c2 + + // Row 2: a2 × [b0,b1,b2,b3] + lo = a2 * b0 + hi = unsignedMultiplyHigh(a2, b0) + prev = w.l2 + s = prev + lo + c1 = if (uLtInline(s, prev)) 1L else 0L + w.l2 = s + carry = hi + c1 + + lo = a2 * b1 + hi = unsignedMultiplyHigh(a2, b1) + prev = w.l3 + s = prev + lo + c1 = if (uLtInline(s, prev)) 1L else 0L + s += carry + c2 = if (uLtInline(s, carry)) 1L else 0L + w.l3 = s + carry = hi + c1 + c2 + + lo = a2 * b2 + hi = unsignedMultiplyHigh(a2, b2) + prev = w.l4 + s = prev + lo + c1 = if (uLtInline(s, prev)) 1L else 0L + s += carry + c2 = if (uLtInline(s, carry)) 1L else 0L + w.l4 = s + carry = hi + c1 + c2 + + lo = a2 * b3 + hi = unsignedMultiplyHigh(a2, b3) + prev = w.l5 + s = prev + lo + c1 = if (uLtInline(s, prev)) 1L else 0L + s += carry + c2 = if (uLtInline(s, carry)) 1L else 0L + w.l5 = s + w.l6 = hi + c1 + c2 + + // Row 3: a3 × [b0,b1,b2,b3] + lo = a3 * b0 + hi = unsignedMultiplyHigh(a3, b0) + prev = w.l3 + s = prev + lo + c1 = if (uLtInline(s, prev)) 1L else 0L + w.l3 = s + carry = hi + c1 + + lo = a3 * b1 + hi = unsignedMultiplyHigh(a3, b1) + prev = w.l4 + s = prev + lo + c1 = if (uLtInline(s, prev)) 1L else 0L + s += carry + c2 = if (uLtInline(s, carry)) 1L else 0L + w.l4 = s + carry = hi + c1 + c2 + + lo = a3 * b2 + hi = unsignedMultiplyHigh(a3, b2) + prev = w.l5 + s = prev + lo + c1 = if (uLtInline(s, prev)) 1L else 0L + s += carry + c2 = if (uLtInline(s, carry)) 1L else 0L + w.l5 = s + carry = hi + c1 + c2 + + lo = a3 * b3 + hi = unsignedMultiplyHigh(a3, b3) + prev = w.l6 + s = prev + lo + c1 = if (uLtInline(s, prev)) 1L else 0L + s += carry + c2 = if (uLtInline(s, carry)) 1L else 0L + w.l6 = s + w.l7 = hi + c1 + c2 + } + + /** Dedicated squaring: w = a² (512-bit result). */ + @Suppress("LongMethod") + fun sqrWide( + w: Wide8, + a: Fe4, + ) { + val a0 = a.l0 + val a1 = a.l1 + val a2 = a.l2 + val a3 = a.l3 + var lo: Long + var hi: Long + var prev: Long + var s: Long + var c1: Long + var c2: Long + var carry: Long + var v: Long + + // Pass 1: cross-products a[i]*a[j] for i < j + w.l0 = 0L + lo = a0 * a1 + w.l1 = lo + carry = unsignedMultiplyHigh(a0, a1) + + lo = a0 * a2 + s = lo + carry + c1 = if (uLtInline(s, lo)) 1L else 0L + w.l2 = s + carry = unsignedMultiplyHigh(a0, a2) + c1 + + lo = a0 * a3 + s = lo + carry + c1 = if (uLtInline(s, lo)) 1L else 0L + w.l3 = s + w.l4 = unsignedMultiplyHigh(a0, a3) + c1 + + lo = a1 * a2 + hi = unsignedMultiplyHigh(a1, a2) + prev = w.l3 + s = prev + lo + c1 = if (uLtInline(s, prev)) 1L else 0L + w.l3 = s + carry = hi + c1 + + lo = a1 * a3 + hi = unsignedMultiplyHigh(a1, a3) + prev = w.l4 + s = prev + lo + c1 = if (uLtInline(s, prev)) 1L else 0L + s += carry + c2 = if (uLtInline(s, carry)) 1L else 0L + w.l4 = s + w.l5 = hi + c1 + c2 + + lo = a2 * a3 + hi = unsignedMultiplyHigh(a2, a3) + prev = w.l5 + s = prev + lo + c1 = if (uLtInline(s, prev)) 1L else 0L + w.l5 = s + w.l6 = hi + c1 + + // Pass 2: double all cross-products (shift left by 1 bit) + v = w.l1 + w.l1 = v shl 1 + var shiftCarry = v ushr 63 + v = w.l2 + w.l2 = (v shl 1) or shiftCarry + shiftCarry = v ushr 63 + v = w.l3 + w.l3 = (v shl 1) or shiftCarry + shiftCarry = v ushr 63 + v = w.l4 + w.l4 = (v shl 1) or shiftCarry + shiftCarry = v ushr 63 + v = w.l5 + w.l5 = (v shl 1) or shiftCarry + shiftCarry = v ushr 63 + v = w.l6 + w.l6 = (v shl 1) or shiftCarry + shiftCarry = v ushr 63 + w.l7 = shiftCarry + + // Pass 3: add diagonal products a[i]² + lo = a0 * a0 + hi = unsignedMultiplyHigh(a0, a0) + w.l0 = lo + s = w.l1 + hi + c1 = if (uLtInline(s, w.l1)) 1L else 0L + w.l1 = s + var dCarry = c1 + + lo = a1 * a1 + hi = unsignedMultiplyHigh(a1, a1) + s = w.l2 + lo + c1 = if (uLtInline(s, w.l2)) 1L else 0L + s += dCarry + c2 = if (uLtInline(s, dCarry)) 1L else 0L + w.l2 = s + prev = w.l3 + hi + val c3a = if (uLtInline(prev, w.l3)) 1L else 0L + prev += c1 + c2 + val c4a = if (uLtInline(prev, c1 + c2)) 1L else 0L + w.l3 = prev + dCarry = c3a + c4a + + lo = a2 * a2 + hi = unsignedMultiplyHigh(a2, a2) + s = w.l4 + lo + c1 = if (uLtInline(s, w.l4)) 1L else 0L + s += dCarry + c2 = if (uLtInline(s, dCarry)) 1L else 0L + w.l4 = s + prev = w.l5 + hi + val c3b = if (uLtInline(prev, w.l5)) 1L else 0L + prev += c1 + c2 + val c4b = if (uLtInline(prev, c1 + c2)) 1L else 0L + w.l5 = prev + dCarry = c3b + c4b + + lo = a3 * a3 + hi = unsignedMultiplyHigh(a3, a3) + s = w.l6 + lo + c1 = if (uLtInline(s, w.l6)) 1L else 0L + s += dCarry + c2 = if (uLtInline(s, dCarry)) 1L else 0L + w.l6 = s + prev = w.l7 + hi + prev += c1 + c2 + w.l7 = prev + } +} + +/** + * Field arithmetic mod p using Fe4 (struct-based, no bounds checks). + */ +internal object Fe4FieldP { + private const val P0 = -4294968273L // 0xFFFFFFFEFFFFFC2F + + fun reduceSelf(a: Fe4) { + if (a.l3 == -1L && a.l2 == -1L && a.l1 == -1L && + (a.l0 xor Long.MIN_VALUE) >= (P0 xor Long.MIN_VALUE) + ) { + a.l0 -= P0 + a.l1 = 0L + a.l2 = 0L + a.l3 = 0L + } + } + + /** Reduce 512-bit value mod p. */ + @Suppress("LongMethod") + fun reduceWide( + out: Fe4, + w: Wide8, + ) { + val c = 4294968273L // 2^32 + 977 + var hcLo: Long + var hcHi: Long + var s1: Long + var s2: Long + var c1: Long + var c2: Long + + // Round 1: acc = lo + hi × C + hcLo = w.l4 * c + hcHi = unsignedMultiplyHigh(w.l4, c) + s1 = w.l0 + hcLo + c1 = if (uLtInline(s1, w.l0)) 1L else 0L + out.l0 = s1 + var carry = hcHi + c1 + + hcLo = w.l5 * c + hcHi = unsignedMultiplyHigh(w.l5, c) + s1 = w.l1 + hcLo + c1 = if (uLtInline(s1, w.l1)) 1L else 0L + s2 = s1 + carry + c2 = if (uLtInline(s2, s1)) 1L else 0L + out.l1 = s2 + carry = hcHi + c1 + c2 + + hcLo = w.l6 * c + hcHi = unsignedMultiplyHigh(w.l6, c) + s1 = w.l2 + hcLo + c1 = if (uLtInline(s1, w.l2)) 1L else 0L + s2 = s1 + carry + c2 = if (uLtInline(s2, s1)) 1L else 0L + out.l2 = s2 + carry = hcHi + c1 + c2 + + hcLo = w.l7 * c + hcHi = unsignedMultiplyHigh(w.l7, c) + s1 = w.l3 + hcLo + c1 = if (uLtInline(s1, w.l3)) 1L else 0L + s2 = s1 + carry + c2 = if (uLtInline(s2, s1)) 1L else 0L + out.l3 = s2 + carry = hcHi + c1 + c2 + + // Round 2: fold carry × C + if (carry != 0L) { + val ccLo = carry * c + val ccHi = unsignedMultiplyHigh(carry, c) + s1 = out.l0 + ccLo + c1 = if (uLtInline(s1, out.l0)) 1L else 0L + out.l0 = s1 + var prop = ccHi + c1 + if (prop != 0L) { + s1 = out.l1 + prop + prop = if (uLtInline(s1, out.l1)) 1L else 0L + out.l1 = s1 + if (prop != 0L) { + s1 = out.l2 + prop + prop = if (uLtInline(s1, out.l2)) 1L else 0L + out.l2 = s1 + if (prop != 0L) { + s1 = out.l3 + prop + prop = if (uLtInline(s1, out.l3)) 1L else 0L + out.l3 = s1 + } + } + } + if (prop != 0L) { + s1 = out.l0 + c + c1 = if (uLtInline(s1, out.l0)) 1L else 0L + out.l0 = s1 + if (c1 != 0L) { + out.l1++ + if (out.l1 == 0L) { + out.l2++ + if (out.l2 == 0L) out.l3++ + } + } + } + } + + reduceSelf(out) + } + + /** out = (a + b) mod p. */ + fun add( + out: Fe4, + a: Fe4, + b: Fe4, + ) { + val carry = Fe4U256.addTo(out, a, b) + if (carry != 0) { + val s1 = out.l0 + 4294968273L + val c1 = if (uLtInline(s1, out.l0)) 1L else 0L + out.l0 = s1 + if (c1 != 0L) { + out.l1++ + if (out.l1 == 0L) { + out.l2++ + if (out.l2 == 0L) out.l3++ + } + } + } + reduceSelf(out) + } + + /** out = (a - b) mod p. */ + fun sub( + out: Fe4, + a: Fe4, + b: Fe4, + ) { + val borrow = Fe4U256.subTo(out, a, b) + if (borrow != 0) { + val s0 = out.l0 + P0 + val c0 = if (uLtInline(s0, out.l0)) 1L else 0L + out.l0 = s0 + if (c0 == 0L) { + if (out.l1 != 0L) { + out.l1-- + } else { + out.l1 = -1L + if (out.l2 != 0L) { + out.l2-- + } else { + out.l2 = -1L + out.l3-- + } + } + } + } + } + + /** out = (a × b) mod p using caller-provided wide buffer. */ + fun mul( + out: Fe4, + a: Fe4, + b: Fe4, + w: Wide8, + ) { + Fe4U256.mulWide(w, a, b) + reduceWide(out, w) + } + + /** out = a² mod p using caller-provided wide buffer. */ + fun sqr( + out: Fe4, + a: Fe4, + w: Wide8, + ) { + Fe4U256.sqrWide(w, a) + reduceWide(out, w) + } +} diff --git a/quartz/src/jvmTest/kotlin/com/vitorpamplona/quartz/utils/secp256k1/Fe4Benchmark.kt b/quartz/src/jvmTest/kotlin/com/vitorpamplona/quartz/utils/secp256k1/Fe4Benchmark.kt new file mode 100644 index 000000000..efae1e5c9 --- /dev/null +++ b/quartz/src/jvmTest/kotlin/com/vitorpamplona/quartz/utils/secp256k1/Fe4Benchmark.kt @@ -0,0 +1,264 @@ +/* + * Copyright (c) 2025 Vitor Pamplona + * + * Permission is hereby granted, free of charge, to any person obtaining a copy of + * this software and associated documentation files (the "Software"), to deal in + * the Software without restriction, including without limitation the rights to use, + * copy, modify, merge, publish, distribute, sublicense, and/or sell copies of the + * Software, and to permit persons to whom the Software is furnished to do so, + * subject to the following conditions: + * + * The above copyright notice and this permission notice shall be included in all + * copies or substantial portions of the Software. + * + * THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR + * IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, FITNESS + * FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE AUTHORS OR + * COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER LIABILITY, WHETHER IN + * AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, OUT OF OR IN CONNECTION + * WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE SOFTWARE. + */ +package com.vitorpamplona.quartz.utils.secp256k1 + +import kotlin.test.Test +import kotlin.test.assertTrue + +/** + * Benchmark comparing LongArray(4) vs Fe4 (struct with 4 @JvmField Long fields) + * for secp256k1 field element operations. + * + * The hypothesis: Fe4 eliminates array bounds checks (laload/lastore → getfield/putfield), + * which should measurably improve performance on platforms where the JIT cannot reliably + * eliminate bounds checks for parameter arrays (Android ART, Kotlin/Native LLVM). + * + * On JVM HotSpot C2, the difference may be smaller because C2 can profile array sizes + * and eliminate constant-index bounds checks. But even on HotSpot, the Fe4 approach + * avoids the array length field load + comparison that precedes each array access. + * + * Run with: ./gradlew :quartz:jvmTest --tests "*.Fe4Benchmark" + */ +class Fe4Benchmark { + // Test vectors: secp256k1 generator point coordinates + private val aBytes = hexToBytes("79BE667EF9DCBBAC55A06295CE870B07029BFCDB2DCE28D959F2815B16F81798") + private val bBytes = hexToBytes("483ADA7726A3C4655DA4FBFC0E1108A8FD17B448A68554199C47D08FFB10D4B8") + + // LongArray representations + private val aArr = U256.fromBytes(aBytes) + private val bArr = U256.fromBytes(bBytes) + private val outArr = LongArray(4) + private val wArr = LongArray(8) + + // Fe4 representations (same values) + private val aFe4 = Fe4(aArr[0], aArr[1], aArr[2], aArr[3]) + private val bFe4 = Fe4(bArr[0], bArr[1], bArr[2], bArr[3]) + private val outFe4 = Fe4() + private val wFe8 = Wide8() + + private data class BenchResult( + val name: String, + val arrayNanos: Long, + val fe4Nanos: Long, + val iterations: Int, + ) { + val arrayNsPerOp get() = arrayNanos / iterations + val fe4NsPerOp get() = fe4Nanos / iterations + val speedup get() = arrayNanos.toDouble() / fe4Nanos.toDouble() + + override fun toString(): String { + val pct = ((speedup - 1.0) * 100).let { if (it >= 0) "+%.1f%%".format(it) else "%.1f%%".format(it) } + return String.format( + " %-20s LongArray: %,8d ns/op Fe4: %,8d ns/op %s", + name, + arrayNsPerOp, + fe4NsPerOp, + pct, + ) + } + } + + private inline fun bench( + name: String, + warmup: Int, + iterations: Int, + crossinline arrayOp: () -> Unit, + crossinline fe4Op: () -> Unit, + ): BenchResult { + // Warmup both generously (C2 compiles at ~10K invocations) + repeat(warmup) { arrayOp() } + repeat(warmup) { fe4Op() } + + // Run 3 rounds, alternating order, take best of each to reduce noise + var bestArr = Long.MAX_VALUE + var bestFe4 = Long.MAX_VALUE + for (round in 0 until 3) { + if (round % 2 == 0) { + // LongArray first + val arrStart = System.nanoTime() + repeat(iterations) { arrayOp() } + bestArr = minOf(bestArr, System.nanoTime() - arrStart) + + val fe4Start = System.nanoTime() + repeat(iterations) { fe4Op() } + bestFe4 = minOf(bestFe4, System.nanoTime() - fe4Start) + } else { + // Fe4 first + val fe4Start = System.nanoTime() + repeat(iterations) { fe4Op() } + bestFe4 = minOf(bestFe4, System.nanoTime() - fe4Start) + + val arrStart = System.nanoTime() + repeat(iterations) { arrayOp() } + bestArr = minOf(bestArr, System.nanoTime() - arrStart) + } + } + + return BenchResult(name, bestArr, bestFe4, iterations) + } + + @Test + fun benchmarkFieldOps() { + // Verify correctness first: both implementations must produce identical results + verifySameResults() + + val results = mutableListOf() + + // --- FieldP.mul (the hottest operation: ~1900 calls/verify) --- + // Use out as input to next iteration to create data dependency chain + aArr.copyInto(outArr) + aFe4.copyFrom(Fe4(aArr[0], aArr[1], aArr[2], aArr[3])) + outFe4.copyFrom(aFe4) + results += + bench( + "FieldP.mul", + 10000, + 200000, + arrayOp = { FieldP.mul(outArr, outArr, bArr, wArr) }, + fe4Op = { Fe4FieldP.mul(outFe4, outFe4, bFe4, wFe8) }, + ) + + // --- FieldP.sqr (second hottest: ~1900 calls/verify) --- + aArr.copyInto(outArr) + outFe4.copyFrom(aFe4) + results += + bench( + "FieldP.sqr", + 10000, + 200000, + arrayOp = { FieldP.sqr(outArr, outArr, wArr) }, + fe4Op = { Fe4FieldP.sqr(outFe4, outFe4, wFe8) }, + ) + + // --- FieldP.add (~750 calls/verify) --- + aArr.copyInto(outArr) + outFe4.copyFrom(aFe4) + results += + bench( + "FieldP.add", + 10000, + 500000, + arrayOp = { FieldP.add(outArr, outArr, bArr) }, + fe4Op = { Fe4FieldP.add(outFe4, outFe4, bFe4) }, + ) + + // --- FieldP.sub (~500 calls/verify) --- + aArr.copyInto(outArr) + outFe4.copyFrom(aFe4) + results += + bench( + "FieldP.sub", + 10000, + 500000, + arrayOp = { FieldP.sub(outArr, outArr, bArr) }, + fe4Op = { Fe4FieldP.sub(outFe4, outFe4, bFe4) }, + ) + + // --- U256.mulWide (raw wide multiply, no reduction) --- + results += + bench( + "U256.mulWide", + 10000, + 200000, + arrayOp = { U256.mulWide(wArr, aArr, bArr) }, + fe4Op = { Fe4U256.mulWide(wFe8, aFe4, bFe4) }, + ) + + // --- U256.sqrWide (raw wide square) --- + results += + bench( + "U256.sqrWide", + 10000, + 200000, + arrayOp = { U256.sqrWide(wArr, aArr) }, + fe4Op = { Fe4U256.sqrWide(wFe8, aFe4) }, + ) + + // Print results + println() + println("=".repeat(80)) + println("Fe4 vs LongArray Benchmark: JVM21/HotSpot C2") + println("=".repeat(80)) + println(" Hypothesis: Fe4 (named fields) eliminates array bounds checks,") + println(" producing faster code especially on ART/Native where JIT is weaker.") + println("-".repeat(80)) + for (r in results) { + println(r) + } + println("=".repeat(80)) + println() + println(" Bytecode analysis:") + println(" LongArray access: aload + iconst + laload (3 insns + implicit bounds check)") + println(" Fe4 field access: aload + getfield (2 insns, no check)") + println() + println(" U256.class: 150 laload/lastore operations (each bounds-checked)") + println(" FieldP.class: 119 laload/lastore operations") + println(" FieldMulFusedKt: 195 laload/lastore operations") + println(" Total: 464 bounds checks in 3 core files") + println() + + // Prevent dead code elimination + assertTrue(outArr[0] != Long.MIN_VALUE || outFe4.l0 != Long.MIN_VALUE) + } + + private fun verifySameResults() { + // Test mul + FieldP.mul(outArr, aArr, bArr, wArr) + Fe4FieldP.mul(outFe4, aFe4, bFe4, wFe8) + assertFe4EqualsArray("mul", outArr, outFe4) + + // Test sqr + FieldP.sqr(outArr, aArr, wArr) + Fe4FieldP.sqr(outFe4, aFe4, wFe8) + assertFe4EqualsArray("sqr", outArr, outFe4) + + // Test add + FieldP.add(outArr, aArr, bArr) + Fe4FieldP.add(outFe4, aFe4, bFe4) + assertFe4EqualsArray("add", outArr, outFe4) + + // Test sub + FieldP.sub(outArr, aArr, bArr) + Fe4FieldP.sub(outFe4, aFe4, bFe4) + assertFe4EqualsArray("sub", outArr, outFe4) + } + + private fun assertFe4EqualsArray( + op: String, + arr: LongArray, + fe: Fe4, + ) { + assertTrue( + arr[0] == fe.l0 && arr[1] == fe.l1 && arr[2] == fe.l2 && arr[3] == fe.l3, + "$op: LongArray[${arr[0]},${arr[1]},${arr[2]},${arr[3]}] != " + + "Fe4[${fe.l0},${fe.l1},${fe.l2},${fe.l3}]", + ) + } + + private fun hexToBytes(hex: String): ByteArray { + val len = hex.length / 2 + val result = ByteArray(len) + for (i in 0 until len) { + result[i] = hex.substring(i * 2, i * 2 + 2).toInt(16).toByte() + } + return result + } +}