diff --git a/quartz/benchmarks/run_all.sh b/quartz/benchmarks/run_all.sh new file mode 100755 index 0000000000..e7bf0f7ba6 --- /dev/null +++ b/quartz/benchmarks/run_all.sh @@ -0,0 +1,302 @@ +#!/usr/bin/env bash +# +# Run all secp256k1 benchmarks and produce a comparison table. +# Run from the amethyst root directory: +# +# ./quartz/benchmarks/run_all.sh +# +# Runs: C native, Kotlin/Native, JVM (always), and Android (if device connected). +# +set -uo pipefail + +ROOT="$(cd "$(dirname "$0")/../.." && pwd)" +BENCH_DIR="$ROOT/quartz/benchmarks" +SO_NAME="libsecp256k1-jni.so" + +# ============================================================================ +# 1. C native benchmark +# ============================================================================ +echo "=== Building C native benchmark ===" + +SO_PATH=$(find "$HOME/.gradle" -path "*/secp256k1-kmp-jni-jvm-linux*" -name "*.jar" 2>/dev/null | head -1) +if [ -z "$SO_PATH" ]; then + echo "ACINQ secp256k1-kmp-jni-jvm-linux JAR not found in Gradle cache." + echo "Run './gradlew :quartz:jvmTest --tests *.Secp256k1Test' first to download it." + exit 1 +fi + +WORK=$(mktemp -d) +trap "rm -rf $WORK" EXIT + +(cd "$WORK" && jar xf "$SO_PATH") 2>/dev/null || true +SO_FILE=$(find "$WORK" -name "$SO_NAME" -path "*linux-x86_64*" 2>/dev/null | head -1) +if [ -z "$SO_FILE" ]; then + echo "Could not find $SO_NAME for linux-x86_64 in JAR." + exit 1 +fi +cp "$SO_FILE" "$WORK/" + +if ! gcc -O2 -o "$WORK/bench" "$BENCH_DIR/secp256k1_native_bench.c" \ + -L"$WORK" -lsecp256k1-jni -Wl,-rpath,"$WORK" 2>&1; then + echo "gcc compilation failed" +fi + +echo "=== Running C native benchmark ===" +C_OUT="" +if [ -f "$WORK/bench" ]; then + C_OUT=$("$WORK/bench" 2>&1) +else + echo "(skipped — compilation failed)" +fi +echo "$C_OUT" +echo "" + +# ============================================================================ +# 2. Kotlin/Native benchmark +# ============================================================================ +echo "=== Running Kotlin/Native benchmark ===" +"$ROOT/gradlew" -p "$ROOT" :quartz:linuxX64Test --tests "*.Secp256k1NativeBenchmark" \ + 2>/dev/null || true + +KN_XML="$ROOT/quartz/build/test-results/linuxX64Test/TEST-linuxX64Test.com.vitorpamplona.quartz.utils.secp256k1.Secp256k1NativeBenchmark.xml" +KN_OUT="" +if [ -f "$KN_XML" ]; then + KN_OUT=$(sed -n '//p' "$KN_XML" | grep -v 'CDATA\|]]>') + echo "$KN_OUT" +else + echo "K/Native benchmark XML not found." +fi +echo "" + +# ============================================================================ +# 3. JVM benchmark +# ============================================================================ +echo "=== Running JVM benchmark ===" +"$ROOT/gradlew" -p "$ROOT" :quartz:jvmTest --tests "*.Secp256k1Benchmark" \ + 2>/dev/null || true + +JVM_XML="$ROOT/quartz/build/test-results/jvmTest/TEST-com.vitorpamplona.quartz.utils.secp256k1.Secp256k1Benchmark.xml" +JVM_OUT="" +if [ -f "$JVM_XML" ]; then + JVM_OUT=$(sed -n '//p' "$JVM_XML" | grep -v 'CDATA\|]]>') + echo "$JVM_OUT" +else + echo "JVM benchmark XML not found." +fi +echo "" + +# ============================================================================ +# 4. Android benchmark (if device/emulator connected) +# ============================================================================ +ANDROID_OUT="" +HAS_ANDROID="" +if command -v adb &>/dev/null && adb devices 2>/dev/null | grep -q "device$"; then + HAS_ANDROID="1" + echo "=== Running Android benchmark (device detected) ===" + "$ROOT/gradlew" -p "$ROOT" :benchmark:connectedBenchmarkAndroidTest \ + -Pandroid.testInstrumentationRunnerArguments.class=com.vitorpamplona.quartz.benchmark.Secp256k1Benchmark \ + 2>/dev/null || true + + # AndroidX Benchmark writes test results XML + ANDROID_XML=$(find "$ROOT/benchmark/build" -path "*test-results*" -name "*.xml" \ + -newer "$ROOT/quartz/benchmarks/run_all.sh" 2>/dev/null | head -1) + if [ -n "$ANDROID_XML" ] && [ -f "$ANDROID_XML" ]; then + ANDROID_OUT=$(cat "$ANDROID_XML") + echo "Android benchmark results found." + else + # Try to find results from the sdcard benchmark output + BENCH_JSON=$(adb shell "ls /sdcard/Download/*benchmark*json 2>/dev/null" 2>/dev/null | head -1 | tr -d '\r') + if [ -n "$BENCH_JSON" ]; then + ANDROID_OUT=$(adb shell "cat '$BENCH_JSON'" 2>/dev/null) + echo "Android benchmark JSON found: $BENCH_JSON" + else + echo "Android benchmark ran but results not found." + fi + fi + echo "" +else + echo "=== Skipping Android benchmark (no device/emulator connected) ===" + echo "" +fi + +# ============================================================================ +# 5. Build comparison table +# ============================================================================ + +# Parse functions +parse_c() { + local op="$1" + [ -z "$op" ] && return + echo "$C_OUT" | grep -E "$op" | head -1 | awk '{print $(NF-1)}' | tr -d ',' +} + +parse_kn() { + local op="$1" + [ -z "$op" ] && return + echo "$KN_OUT" | grep -E "$op" | head -1 | awk '{print $(NF-1)}' | tr -d ',' +} + +parse_jvm_jni() { + local op="$1" + [ -z "$op" ] && return + echo "$JVM_OUT" | grep -E "$op" | head -1 | \ + sed 's/.*Native:[[:space:]]*//' | awk '{print $1}' | tr -d ',' +} + +parse_jvm_kotlin() { + local op="$1" + [ -z "$op" ] && return + echo "$JVM_OUT" | grep -E "$op" | head -1 | \ + sed 's/.*Kotlin:[[:space:]]*//' | awk '{print $1}' | tr -d ',' +} + +# Android benchmark: parse ns from XML test results +# Format: +# AndroidX Benchmark puts the median ns in the test output +parse_android() { + local op="$1" + [ -z "$op" ] || [ -z "$ANDROID_OUT" ] && return + # Try to extract ns/op from the benchmark output and convert to ops/sec + local ns + ns=$(echo "$ANDROID_OUT" | grep -i "${op}" | grep -oP '[\d,]+(?=\s*ns)' | head -1 | tr -d ',') + if [ -n "$ns" ] && [ "$ns" != "0" ]; then + awk "BEGIN { printf \"%d\", 1000000000 / $ns }" + fi +} + +fmt() { + local v="$1" + if [ -z "$v" ] || [ "$v" = "--" ]; then + echo "--" + else + echo "$v" | sed ':a;s/\B[0-9]\{3\}\>/,&/;ta' + fi +} + +ratio() { + local a="$1" b="$2" + if [ -n "$a" ] && [ -n "$b" ] && [ "$b" != "0" ] && [ "$a" != "0" ]; then + awk "BEGIN { printf \"%.1fx\", $a / $b }" + else + echo "—" + fi +} + +# Operation definitions: label|c_pattern|kn_pattern|jvm_pattern|android_pattern|quartz_only +# quartz_only=1: no C/JNI equivalent, blank those columns +OPS=( + "verifySchnorr|verifySchnorr[[:space:]]|verifySchnorr[[:space:]]|verifySchnorr[[:space:]]|verifySchnorrOurs|" + "verifySchnorrFast||verifySchnorrFast|verifySchnorrFast|verifySchnorrFastOurs|1" + "signSchnorr|signSchnorr[[:space:]]|signSchnorr[[:space:]]|signSchnorr[[:space:]]|signSchnorrOurs|" + "signSchnorr (cached)||signSchnorr .cached|signSchnorr .cached|signSchnorrCachedPkOurs|1" + "compressedPubKeyFor|compressedPubKeyFor|compressedPubKeyFor|compressedPubKeyFor|compressedPubKeyForOurs|" + "secKeyVerify|secKeyVerify|secKeyVerify|secKeyVerify|secKeyVerifyOurs|" + "privKeyTweakAdd|privKeyTweakAdd|privKeyTweakAdd|privKeyTweakAdd|privateKeyAddOurs|" + "ecdh/tweakMul|ecPubKeyTweakMul|ecdhXOnly|ecdhXOnly|ecdhXOnlyOurs|" +) + +# Determine column count based on Android availability +if [ -n "$HAS_ANDROID" ]; then + COL_FMT="%-24s %14s %14s %14s %14s %14s\n" + SEP_FMT="%-24s %14s %14s %14s %14s %14s\n" +else + COL_FMT="%-24s %14s %14s %14s %14s\n" + SEP_FMT="%-24s %14s %14s %14s %14s\n" +fi + +echo "" +echo "==============================================================================================" +echo "COMPARISON TABLE — ops/sec (higher is better)" +echo "==============================================================================================" +echo "" + +if [ -n "$HAS_ANDROID" ]; then + printf "$COL_FMT" "" "libsecp256k1" "libsecp256k1" "Quartz" "Quartz" "Quartz" + printf "$COL_FMT" "Operation" "C (no JVM)" "JVM (C+JNI)" "JVM Kotlin" "K/Native" "Android" + printf "$SEP_FMT" "——————————————————————————" "——————————————" "——————————————" "——————————————" "——————————————" "——————————————" +else + printf "$COL_FMT" "" "libsecp256k1" "libsecp256k1" "Quartz" "Quartz" + printf "$COL_FMT" "Operation" "C (no JVM)" "JVM (C+JNI)" "JVM Kotlin" "K/Native" + printf "$SEP_FMT" "——————————————————————————" "——————————————" "——————————————" "——————————————" "——————————————" +fi + +for entry in "${OPS[@]}"; do + IFS='|' read -r label c_pat kn_pat jvm_pat android_pat qonly <<< "$entry" + + c_ops=$(parse_c "$c_pat") + kn_ops=$(parse_kn "$kn_pat") + jvm_jni_ops=$(parse_jvm_jni "$jvm_pat") + jvm_k_ops=$(parse_jvm_kotlin "$jvm_pat") + android_ops=$(parse_android "$android_pat") + + # Quartz-only: no C/JNI equivalent + [ "$qonly" = "1" ] && c_ops="" && jvm_jni_ops="" + + if [ -n "$HAS_ANDROID" ]; then + printf "$COL_FMT" \ + "$label" \ + "$(fmt "${c_ops:---}")" \ + "$(fmt "${jvm_jni_ops:---}")" \ + "$(fmt "${jvm_k_ops:---}")" \ + "$(fmt "${kn_ops:---}")" \ + "$(fmt "${android_ops:---}")" + else + printf "$COL_FMT" \ + "$label" \ + "$(fmt "${c_ops:---}")" \ + "$(fmt "${jvm_jni_ops:---}")" \ + "$(fmt "${jvm_k_ops:---}")" \ + "$(fmt "${kn_ops:---}")" + fi +done + +echo "" +echo "——————————————————————————————————————————————————————————————————————————————————————————————" +echo "" +echo "Ratios — lower is closer to native (1.0x = parity):" +echo "" + +if [ -n "$HAS_ANDROID" ]; then + printf "%-24s %14s %14s %14s\n" \ + "" "C (no JVM) vs" "JVM (C+JNI) vs" "libsecp C+JNI vs" + printf "%-24s %14s %14s %14s\n" \ + "Operation" "K/Native" "JVM Kotlin" "Android" + printf "%-24s %14s %14s %14s\n" \ + "——————————————————————————" "——————————————" "——————————————" "——————————————" +else + printf "%-24s %14s %14s\n" \ + "" "C (no JVM) vs" "JVM (C+JNI) vs" + printf "%-24s %14s %14s\n" \ + "Operation" "K/Native" "JVM Kotlin" + printf "%-24s %14s %14s\n" \ + "——————————————————————————" "——————————————" "——————————————" +fi + +for entry in "${OPS[@]}"; do + IFS='|' read -r label c_pat kn_pat jvm_pat android_pat qonly <<< "$entry" + + # Skip Quartz-only in ratio table + [ "$qonly" = "1" ] && continue + + c_ops=$(parse_c "$c_pat") + kn_ops=$(parse_kn "$kn_pat") + jvm_jni_ops=$(parse_jvm_jni "$jvm_pat") + jvm_k_ops=$(parse_jvm_kotlin "$jvm_pat") + android_ops=$(parse_android "$android_pat") + + r_c_kn=$(ratio "${c_ops:-0}" "${kn_ops:-0}") + r_jni_jvm=$(ratio "${jvm_jni_ops:-0}" "${jvm_k_ops:-0}") + + if [ -n "$HAS_ANDROID" ]; then + # For Android ratio, compare native C JNI benchmark on same device + # We don't have a separate C-on-Android number, so use the JNI native from JVM + # as an approximation. Better: parse the native Android benchmark results. + r_android=$(ratio "${jvm_jni_ops:-0}" "${android_ops:-0}") + printf "%-24s %14s %14s %14s\n" "$label" "$r_c_kn" "$r_jni_jvm" "$r_android" + else + printf "%-24s %14s %14s\n" "$label" "$r_c_kn" "$r_jni_jvm" + fi +done + +echo "" +echo "==============================================================================================" diff --git a/quartz/benchmarks/secp256k1_native_bench.c b/quartz/benchmarks/secp256k1_native_bench.c new file mode 100644 index 0000000000..f98d748869 --- /dev/null +++ b/quartz/benchmarks/secp256k1_native_bench.c @@ -0,0 +1,106 @@ +// Standalone C benchmark for libsecp256k1 -- no JVM, no JNI, no ART. +// Links against the ACINQ secp256k1-kmp-jni .so to benchmark raw C performance. +// Uses the same test vectors as the Kotlin benchmarks for direct comparison. +// +// Build: +// Extract the native .so from the ACINQ JAR, then: +// gcc -O2 -o bench secp256k1_native_bench.c -L. -lsecp256k1-jni -Wl,-rpath,. +// ./bench +#include +#include +#include +#include +#include + +typedef struct { unsigned char data[64]; } secp256k1_pubkey; +typedef struct { unsigned char data[64]; } secp256k1_xonly_pubkey; +typedef struct { unsigned char data[96]; } secp256k1_keypair; +typedef struct secp256k1_context_struct secp256k1_context; + +/* flags: SIGN=0x201, VERIFY=0x101 */ +#define SECP256K1_FLAGS (0x201 | 0x101) +#define SECP256K1_EC_COMPRESSED 258 + +extern secp256k1_context *secp256k1_context_create(unsigned int flags); +extern void secp256k1_context_destroy(secp256k1_context *ctx); +extern int secp256k1_ec_seckey_verify(const secp256k1_context *ctx, const unsigned char *seckey); +extern int secp256k1_ec_pubkey_create(const secp256k1_context *ctx, secp256k1_pubkey *pubkey, const unsigned char *seckey); +extern int secp256k1_ec_pubkey_serialize(const secp256k1_context *ctx, unsigned char *output, size_t *outputlen, const secp256k1_pubkey *pubkey, unsigned int flags); +extern int secp256k1_schnorrsig_sign32(const secp256k1_context *ctx, unsigned char *sig64, const unsigned char *msg32, const secp256k1_keypair *keypair, const unsigned char *aux_rand32); +extern int secp256k1_schnorrsig_verify(const secp256k1_context *ctx, const unsigned char *sig64, const unsigned char *msg, size_t msglen, const secp256k1_xonly_pubkey *pubkey); +extern int secp256k1_keypair_create(const secp256k1_context *ctx, secp256k1_keypair *keypair, const unsigned char *seckey); +extern int secp256k1_keypair_xonly_pub(const secp256k1_context *ctx, secp256k1_xonly_pubkey *pubkey, int *pk_parity, const secp256k1_keypair *keypair); +extern int secp256k1_ec_seckey_tweak_add(const secp256k1_context *ctx, unsigned char *seckey, const unsigned char *tweak); +extern int secp256k1_ec_pubkey_tweak_mul(const secp256k1_context *ctx, secp256k1_pubkey *pubkey, const unsigned char *tweak32); + +static void hex2bin(const char *hex, unsigned char *out, size_t len) { + for (size_t i = 0; i < len; i++) { unsigned v; sscanf(hex+2*i, "%2x", &v); out[i] = v; } +} + +static uint64_t now_ns(void) { + struct timespec ts; clock_gettime(CLOCK_MONOTONIC, &ts); + return (uint64_t)ts.tv_sec * 1000000000ULL + ts.tv_nsec; +} + +#define BENCH(name, warmup, iters, body) do { \ + for (int _w = 0; _w < (warmup); _w++) { body; } \ + uint64_t _start = now_ns(); \ + for (int _i = 0; _i < (iters); _i++) { body; } \ + uint64_t _el = now_ns() - _start; \ + printf(" %-24s %8llu ns/op %8llu ops/s\n", name, \ + (unsigned long long)(_el / (iters)), \ + (unsigned long long)((uint64_t)(iters) * 1000000000ULL / _el)); \ +} while(0) + +int main(void) { + secp256k1_context *ctx = secp256k1_context_create(SECP256K1_FLAGS); + unsigned char priv[32], msg[32], aux[32], priv2[32]; + hex2bin("67E56582298859DDAE725F972992A07C6C4FB9F62A8FFF58CE3CA926A1063530", priv, 32); + hex2bin("243F6A8885A308D313198A2E03707344A4093822299F31D0082EFA98EC4E6C89", msg, 32); + hex2bin("0000000000000000000000000000000000000000000000000000000000000001", aux, 32); + hex2bin("3982F19BEF1615BCCFBB05E321C10E1D4CBA3DF0E841C2E41EEB6016347653C3", priv2, 32); + + secp256k1_keypair kp; secp256k1_keypair_create(ctx, &kp, priv); + secp256k1_xonly_pubkey xpub; secp256k1_keypair_xonly_pub(ctx, &xpub, NULL, &kp); + secp256k1_pubkey pub; secp256k1_ec_pubkey_create(ctx, &pub, priv); + unsigned char sig[64]; secp256k1_schnorrsig_sign32(ctx, sig, msg, &kp, aux); + + if (!secp256k1_schnorrsig_verify(ctx, sig, msg, 32, &xpub)) { + fprintf(stderr, "verify failed!\n"); return 1; + } + + volatile int r; + printf("================================================================================\n"); + printf("secp256k1 Benchmark: C libsecp256k1 (direct, no JNI/JVM) on x86_64\n"); + printf("================================================================================\n"); + + BENCH("verifySchnorr", 2000, 5000, + r = secp256k1_schnorrsig_verify(ctx, sig, msg, 32, &xpub)); + + BENCH("signSchnorr", 1000, 3000, + secp256k1_schnorrsig_sign32(ctx, sig, msg, &kp, aux)); + + { unsigned char c[33]; size_t cl; + BENCH("compressedPubKeyFor", 1000, 5000, { + secp256k1_ec_pubkey_create(ctx, &pub, priv); + cl = 33; secp256k1_ec_pubkey_serialize(ctx, c, &cl, &pub, SECP256K1_EC_COMPRESSED); + }); } + + BENCH("secKeyVerify", 5000, 200000, + r = secp256k1_ec_seckey_verify(ctx, priv)); + + { unsigned char tw[32]; + BENCH("privKeyTweakAdd", 1000, 50000, { + memcpy(tw, priv, 32); + secp256k1_ec_seckey_tweak_add(ctx, tw, priv2); + }); } + + { secp256k1_pubkey p2; secp256k1_ec_pubkey_create(ctx, &p2, priv2); + BENCH("ecPubKeyTweakMul", 1000, 3000, + secp256k1_ec_pubkey_tweak_mul(ctx, &p2, priv)); + } + + printf("================================================================================\n"); + secp256k1_context_destroy(ctx); + return 0; +} diff --git a/quartz/src/androidMain/kotlin/com/vitorpamplona/quartz/utils/secp256k1/FieldMulPlatform.android.kt b/quartz/src/androidMain/kotlin/com/vitorpamplona/quartz/utils/secp256k1/FieldMulPlatform.android.kt index f25f1e0162..14f9bda718 100644 --- a/quartz/src/androidMain/kotlin/com/vitorpamplona/quartz/utils/secp256k1/FieldMulPlatform.android.kt +++ b/quartz/src/androidMain/kotlin/com/vitorpamplona/quartz/utils/secp256k1/FieldMulPlatform.android.kt @@ -21,26 +21,26 @@ package com.vitorpamplona.quartz.utils.secp256k1 /** - * Android field multiply/square — uses pure-Kotlin fallback for multiply-high. + * Uses the pure-Kotlin fused fallback for all API levels. * - * WHY NOT use Math.unsignedMultiplyHigh (API 35+) or Math.multiplyHigh (API 31+)? - * Because D8/R8 desugaring generates a synthetic backport wrapper - * (ExternalSyntheticBackport0.m) for ANY Math.xxx call when minSdk < the API - * that introduced it (minSdk=26 < 31/35). This backport wrapper: - * - Adds 139ns per call (traced on Pixel 8 with Android 16) - * - Is called 24,875 times per Schnorr verify - * - Costs 3.45ms total = 17.5% of verify time - * The pure-Kotlin fallback (4 Long multiplies + shifts) avoids the backport - * entirely and is FASTER than the backported intrinsic path. + * THREE approaches to reach hardware UMULH were tested and all failed: * - * WHY NOT use API-level dispatch (fieldMulApi35/Api31/Fallback)? - * Profiling showed the D8 backport overhead dominates. The UMULH/SMULH - * intrinsic behind the backport is ~1ns, but the backport wrapper adds ~138ns. - * The pure fallback at ~10-20ns is much faster than 1+138=139ns. + * 1. Direct Math.unsignedMultiplyHigh: D8 replaces with pure-Java backport + * when app minSdk=26 < 35, even inside libraries with minSdk=35. * - * ALSO: uLt calls accounted for 20.7% of verify time (49,924 calls × 82ns each) - * because it's an expect/actual function (can't be inline). The inline XOR - * comparison in the crossinline lambda eliminates all those function calls. + * 2. MethodHandle.invokeExact (crypto-intrinsics Java module): Bytecode correct + * (invoke-polymorphic JJ→J, zero boxing), but ART can't inline through + * invoke-polymorphic. 25ns/call × 12K calls = 2.2x regression. + * + * 3. Full fieldMulReduce in Java with minSdk=35 module: D8 desugaring runs + * at APP level with app's minSdk=26, not library's minSdk=35. All + * Math.unsignedMultiplyHigh calls still get backported. Plus Java's + * Long.compareUnsigned (282ns/call) is slower than Kotlin's inlined XOR trick. + * + * The Kotlin inline+crossinline pattern produces the best ART code: + * - unsignedMultiplyHighFallback is inlined at each call site (zero dispatch) + * - uLtInline uses XOR+compare (zero method call, ~0ns overhead) + * - The fused function stays within ART's inlining budget */ internal actual fun fieldMulReduce( out: LongArray, diff --git a/quartz/src/commonMain/kotlin/com/vitorpamplona/quartz/utils/secp256k1/ECPoint.kt b/quartz/src/commonMain/kotlin/com/vitorpamplona/quartz/utils/secp256k1/ECPoint.kt index bb31fb6455..c9279fd1c0 100644 --- a/quartz/src/commonMain/kotlin/com/vitorpamplona/quartz/utils/secp256k1/ECPoint.kt +++ b/quartz/src/commonMain/kotlin/com/vitorpamplona/quartz/utils/secp256k1/ECPoint.kt @@ -68,6 +68,14 @@ internal object ECPoint { /** Curve constant b = 7 in y² = x³ + 7. */ private val B = longArrayOf(7L, 0L, 0L, 0L) + // Thread-local scratch — declared before precomputed tables because + // buildGOddTable() and buildCombTable() use doublePoint/addPoints which + // call scratch.get(). Must be initialized before those table fields. + private val scratch = ScratchLocal { PointScratch() } + + /** Get thread-local scratch. Call once at the top-level entry point. */ + internal fun getScratch(): PointScratch = scratch.get() + /** * wNAF window width for the G-side of scalar multiplication (mulDoubleG/verify). * Width w uses a table of 2^(w-2) odd multiples. Larger windows mean fewer @@ -84,14 +92,14 @@ internal object ECPoint { /** * Precomputed G odd-multiples for wNAF: gOddTable[i] = (2i+1)·G as affine, for i in 0..G_TABLE_SIZE-1. - * Used by mulG and mulDoubleG. Lazily initialized on first use. + * Used by mulG and mulDoubleG. Eagerly initialized to avoid SynchronizedLazyImpl + * dispatch on every verify call (~200 instructions of lock-check overhead per access). */ - private val gOddTable: Array by lazy { buildGOddTable() } + private val gOddTable: Array = buildGOddTable() /** Precomputed λ(G) odd-multiples for GLV: gLamTable[i] = λ((2i+1)·G) as affine. */ - private val gLamTable: Array by lazy { + private val gLamTable: Array = Array(G_TABLE_SIZE) { AffinePoint(FieldP.mul(gOddTable[it].x, Glv.BETA), gOddTable[it].y.copyOf()) } - } // ==================== P-side wNAF table cache ==================== // @@ -146,7 +154,7 @@ internal object ECPoint { // Step 1: compute prefix products of Z coordinates // prods[i] = z[0] * z[1] * ... * z[i] val prods = Array(n) { LongArray(4) } - jac[0].z.copyInto(prods[0]) + jac[0].z.copyInto(prods[0], 0, 0, 4) for (i in 1 until n) { FieldP.mul(prods[i], prods[i - 1], jac[i].z) } @@ -203,7 +211,7 @@ internal object ECPoint { private const val COMB_SPACING = 4 private const val COMB_POINTS = 1 shl COMB_TEETH // 64 - private val combTable: Array by lazy { buildCombTable() } + private val combTable: Array = buildCombTable() private fun buildCombTable(): Array { // Tooth base points: toothG[i] = 2^(i * SPACING) * G @@ -247,13 +255,6 @@ internal object ECPoint { } } - // ==================== Thread-local scratch ==================== - - private val scratch = ScratchLocal { PointScratch() } - - /** Get thread-local scratch. Call once at the top-level entry point. */ - internal fun getScratch(): PointScratch = scratch.get() - // ==================== Point Doubling (3M + 4S) ==================== // Point doubling: out = 2·p. @@ -454,13 +455,12 @@ internal object ECPoint { out: MutablePoint, p: MutablePoint, scalar: LongArray, + s: PointScratch = scratch.get(), ) { if (U256.isZero(scalar) || p.isInfinity()) { out.setInfinity() return } - - val s = scratch.get() val wnd = 5 val tableSize = 1 shl (wnd - 2) // 8 entries @@ -482,8 +482,8 @@ internal object ECPoint { val pLamOddJac = s.pLamOddJac for (i in 0 until tableSize) { FieldP.mul(pLamOddJac[i].x, pOddJac[i].x, Glv.BETA, s.w) - pOddJac[i].y.copyInto(pLamOddJac[i].y) - pOddJac[i].z.copyInto(pLamOddJac[i].z) + pOddJac[i].y.copyInto(pLamOddJac[i].y, 0, 0, 4) + pOddJac[i].z.copyInto(pLamOddJac[i].z, 0, 0, 4) } // Effective-affine: batch-convert with shared Z inversion @@ -537,13 +537,12 @@ internal object ECPoint { fun mulG( out: MutablePoint, scalar: LongArray, + s: PointScratch = scratch.get(), ) { if (U256.isZero(scalar)) { out.setInfinity() return } - - val s = scratch.get() val table = combTable // Ping-pong: alternate between out and s.mixTmp to avoid copyFrom after @@ -597,8 +596,8 @@ internal object ECPoint { s: LongArray, p: MutablePoint, e: LongArray, + sc: PointScratch = scratch.get(), ) { - val sc = scratch.get() val wP = 5 // Window for P-side (table built per-call, keep small) val pTableSize = 1 shl (wP - 2) // 8 entries for P @@ -640,8 +639,8 @@ internal object ECPoint { val pLamOddJac = sc.pLamOddJac for (i in 0 until pTableSize) { FieldP.mul(pLamOddJac[i].x, pOddJac[i].x, Glv.BETA, sc.w) - pOddJac[i].y.copyInto(pLamOddJac[i].y) - pOddJac[i].z.copyInto(pLamOddJac[i].z) + pOddJac[i].y.copyInto(pLamOddJac[i].y, 0, 0, 4) + pOddJac[i].z.copyInto(pLamOddJac[i].z, 0, 0, 4) } // Batch-convert to affine (into scratch arrays) batchToAffinePair(pOddJac, pLamOddJac, sc.pOddAff, sc.pLamOddAff, sc) @@ -871,7 +870,7 @@ internal object ECPoint { // Build prefix products of Z (shared between a and b) val cumZ = s.cumZ - a[0].z.copyInto(cumZ[0]) + a[0].z.copyInto(cumZ[0], 0, 0, 4) for (i in 1 until n) { FieldP.mul(cumZ[i], cumZ[i - 1], a[i].z, w) } diff --git a/quartz/src/commonMain/kotlin/com/vitorpamplona/quartz/utils/secp256k1/FieldMulFused.kt b/quartz/src/commonMain/kotlin/com/vitorpamplona/quartz/utils/secp256k1/FieldMulFused.kt new file mode 100644 index 0000000000..4d3362e2a2 --- /dev/null +++ b/quartz/src/commonMain/kotlin/com/vitorpamplona/quartz/utils/secp256k1/FieldMulFused.kt @@ -0,0 +1,389 @@ +/* + * 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 + +// ===================================================================================== +// FUSED 128-BIT MULTIPLY: computes both lo and hi from shared sub-products. +// +// The current fieldMulReduceWith uses: +// lo = a * b (1 hardware mul — CPU computes all sub-products internally) +// hi = umulh(a, b) (4 imul fallback — recomputes the same 4 sub-products) +// +// This fused version computes both from 4 shared sub-products (aLo*bLo, aLo*bHi, +// aHi*bLo, aHi*bHi), saving 1 multiply per wide product. +// +// With ~20 wide multiplies per field mul, that's ~20 fewer multiplies per field op. +// Across ~750 field ops per verify: ~15,000 fewer multiply instructions. +// +// The inline lambda consumer pattern returns two values without heap allocation: +// mulFull(a, b) { lo, hi -> ... } // both inlined at call site +// +// USE ON: platforms where umulh is a software fallback (Android <35, K/Native). +// DO NOT USE ON: JVM/HotSpot where Math.unsignedMultiplyHigh is a hardware intrinsic. +// ===================================================================================== + +private const val MASK32 = 0xFFFFFFFFL + +/** + * Compute full 64×64 → 128-bit unsigned multiply from 4 sub-products. + * Passes both lo and hi to [consume] which is inlined at the call site. + */ +@Suppress("NOTHING_TO_INLINE") +private inline fun mulFull( + a: Long, + b: Long, + consume: (lo: Long, hi: Long) -> Unit, +) { + val aLo = a and MASK32 + val aHi = a ushr 32 + val bLo = b and MASK32 + val bHi = b ushr 32 + val ll = aLo * bLo + val lh = aLo * bHi + val hl = aHi * bLo + val hh = aHi * bHi + val midSum = (ll ushr 32) + (hl and MASK32) + (lh and MASK32) + consume( + (ll and MASK32) or (midSum shl 32), + hh + (hl ushr 32) + (lh ushr 32) + (midSum ushr 32), + ) +} + +/** Unsigned multiply high only, from 4 sub-products. Used in reduction stage. */ +@Suppress("NOTHING_TO_INLINE") +private inline fun umulhFused( + a: Long, + b: Long, +): Long { + val aLo = a and MASK32 + val aHi = a ushr 32 + val bLo = b and MASK32 + val bHi = b ushr 32 + val mid1 = aHi * bLo + val mid2 = aLo * bHi + val low = aLo * bLo + val carry = ((low ushr 32) + (mid1 and MASK32) + (mid2 and MASK32)) ushr 32 + return (aHi * bHi) + (mid1 ushr 32) + (mid2 ushr 32) + carry +} + +/** + * Fused field multiply: out = (a × b) mod p. + * Uses shared sub-products for each 128-bit multiply (4 imul instead of 5). + */ +@Suppress("LongMethod") +internal fun fieldMulReduceFused( + out: LongArray, + a: LongArray, + b: LongArray, + w: LongArray, +) { + val a0 = a[0] + val a1 = a[1] + val a2 = a[2] + val a3 = a[3] + val b0 = b[0] + val b1 = b[1] + val b2 = b[2] + val b3 = b[3] + var s: Long = 0L + var c1: Long + var c2: Long + var carry: Long = 0L + + // Row 0: a0 × [b0,b1,b2,b3] + mulFull(a0, b0) { lo, hi -> + w[0] = lo + carry = hi + } + + mulFull(a0, b1) { lo, hi -> + s = lo + carry + c1 = if (uLtInline(s, lo)) 1L else 0L + w[1] = s + carry = hi + c1 + } + mulFull(a0, b2) { lo, hi -> + s = lo + carry + c1 = if (uLtInline(s, lo)) 1L else 0L + w[2] = s + carry = hi + c1 + } + mulFull(a0, b3) { lo, hi -> + s = lo + carry + c1 = if (uLtInline(s, lo)) 1L else 0L + w[3] = s + w[4] = hi + c1 + } + + // Row 1: a1 × [b0,b1,b2,b3] + mulFull(a1, b0) { lo, hi -> + val prev = w[1] + s = prev + lo + c1 = if (uLtInline(s, prev)) 1L else 0L + w[1] = s + carry = hi + c1 + } + mulFull(a1, b1) { lo, hi -> + val prev = w[2] + s = prev + lo + c1 = if (uLtInline(s, prev)) 1L else 0L + s += carry + c2 = if (uLtInline(s, carry)) 1L else 0L + w[2] = s + carry = hi + c1 + c2 + } + mulFull(a1, b2) { lo, hi -> + val prev = w[3] + s = prev + lo + c1 = if (uLtInline(s, prev)) 1L else 0L + s += carry + c2 = if (uLtInline(s, carry)) 1L else 0L + w[3] = s + carry = hi + c1 + c2 + } + mulFull(a1, b3) { lo, hi -> + val prev = w[4] + s = prev + lo + c1 = if (uLtInline(s, prev)) 1L else 0L + s += carry + c2 = if (uLtInline(s, carry)) 1L else 0L + w[4] = s + w[5] = hi + c1 + c2 + } + + // Row 2: a2 × [b0,b1,b2,b3] + mulFull(a2, b0) { lo, hi -> + val prev = w[2] + s = prev + lo + c1 = if (uLtInline(s, prev)) 1L else 0L + w[2] = s + carry = hi + c1 + } + mulFull(a2, b1) { lo, hi -> + val prev = w[3] + s = prev + lo + c1 = if (uLtInline(s, prev)) 1L else 0L + s += carry + c2 = if (uLtInline(s, carry)) 1L else 0L + w[3] = s + carry = hi + c1 + c2 + } + mulFull(a2, b2) { lo, hi -> + val prev = w[4] + s = prev + lo + c1 = if (uLtInline(s, prev)) 1L else 0L + s += carry + c2 = if (uLtInline(s, carry)) 1L else 0L + w[4] = s + carry = hi + c1 + c2 + } + mulFull(a2, b3) { lo, hi -> + val prev = w[5] + s = prev + lo + c1 = if (uLtInline(s, prev)) 1L else 0L + s += carry + c2 = if (uLtInline(s, carry)) 1L else 0L + w[5] = s + w[6] = hi + c1 + c2 + } + + // Row 3: a3 × [b0,b1,b2,b3] + mulFull(a3, b0) { lo, hi -> + val prev = w[3] + s = prev + lo + c1 = if (uLtInline(s, prev)) 1L else 0L + w[3] = s + carry = hi + c1 + } + mulFull(a3, b1) { lo, hi -> + val prev = w[4] + s = prev + lo + c1 = if (uLtInline(s, prev)) 1L else 0L + s += carry + c2 = if (uLtInline(s, carry)) 1L else 0L + w[4] = s + carry = hi + c1 + c2 + } + mulFull(a3, b2) { lo, hi -> + val prev = w[5] + s = prev + lo + c1 = if (uLtInline(s, prev)) 1L else 0L + s += carry + c2 = if (uLtInline(s, carry)) 1L else 0L + w[5] = s + carry = hi + c1 + c2 + } + mulFull(a3, b3) { lo, hi -> + val prev = w[6] + s = prev + lo + c1 = if (uLtInline(s, prev)) 1L else 0L + s += carry + c2 = if (uLtInline(s, carry)) 1L else 0L + w[6] = s + w[7] = hi + c1 + c2 + } + + // Reduction uses umulh only (lo computed with hardware *) + reduceWideInline(out, w) { x, y -> umulhFused(x, y) } +} + +/** + * Fused field squaring: out = a² mod p. + * Uses shared sub-products for each 128-bit multiply. + */ +@Suppress("LongMethod") +internal fun fieldSqrReduceFused( + out: LongArray, + a: LongArray, + w: LongArray, +) { + val a0 = a[0] + val a1 = a[1] + val a2 = a[2] + val a3 = a[3] + var s: Long = 0L + var c1: Long + var c2: Long + var carry: Long = 0L + + // Pass 1: cross-products a[i]*a[j] for i < j + w[0] = 0L + mulFull(a0, a1) { lo, hi -> + w[1] = lo + carry = hi + } + + mulFull(a0, a2) { lo, hi -> + s = lo + carry + c1 = if (uLtInline(s, lo)) 1L else 0L + w[2] = s + carry = hi + c1 + } + mulFull(a0, a3) { lo, hi -> + s = lo + carry + c1 = if (uLtInline(s, lo)) 1L else 0L + w[3] = s + w[4] = hi + c1 + } + + mulFull(a1, a2) { lo, hi -> + val prev = w[3] + s = prev + lo + c1 = if (uLtInline(s, prev)) 1L else 0L + w[3] = s + carry = hi + c1 + } + mulFull(a1, a3) { lo, hi -> + val prev = w[4] + s = prev + lo + c1 = if (uLtInline(s, prev)) 1L else 0L + s += carry + c2 = if (uLtInline(s, carry)) 1L else 0L + w[4] = s + w[5] = hi + c1 + c2 + } + + mulFull(a2, a3) { lo, hi -> + val prev = w[5] + s = prev + lo + c1 = if (uLtInline(s, prev)) 1L else 0L + w[5] = s + w[6] = hi + c1 + } + + // Pass 2: double all cross-products (shift left by 1 bit) + var v = w[1] + w[1] = v shl 1 + var shiftCarry = v ushr 63 + v = w[2] + w[2] = (v shl 1) or shiftCarry + shiftCarry = v ushr 63 + v = w[3] + w[3] = (v shl 1) or shiftCarry + shiftCarry = v ushr 63 + v = w[4] + w[4] = (v shl 1) or shiftCarry + shiftCarry = v ushr 63 + v = w[5] + w[5] = (v shl 1) or shiftCarry + shiftCarry = v ushr 63 + v = w[6] + w[6] = (v shl 1) or shiftCarry + shiftCarry = v ushr 63 + w[7] = shiftCarry + + // Pass 3: add diagonal products a[i]² + // Use the original pattern: lo = a*a, hi = umulh(a,a) since the diagonal + // products need careful carry threading that's hard to do inside lambdas. + var dLo: Long + var dHi: Long + var prev: Long + + dLo = a0 * a0 + dHi = umulhFused(a0, a0) + w[0] = dLo + s = w[1] + dHi + c1 = if (uLtInline(s, w[1])) 1L else 0L + w[1] = s + var dCarry = c1 + + dLo = a1 * a1 + dHi = umulhFused(a1, a1) + s = w[2] + dLo + c1 = if (uLtInline(s, w[2])) 1L else 0L + s += dCarry + c2 = if (uLtInline(s, dCarry)) 1L else 0L + w[2] = s + prev = w[3] + dHi + val c3a = if (uLtInline(prev, w[3])) 1L else 0L + prev += c1 + c2 + val c4a = if (uLtInline(prev, c1 + c2)) 1L else 0L + w[3] = prev + dCarry = c3a + c4a + + dLo = a2 * a2 + dHi = umulhFused(a2, a2) + s = w[4] + dLo + c1 = if (uLtInline(s, w[4])) 1L else 0L + s += dCarry + c2 = if (uLtInline(s, dCarry)) 1L else 0L + w[4] = s + prev = w[5] + dHi + val c3b = if (uLtInline(prev, w[5])) 1L else 0L + prev += c1 + c2 + val c4b = if (uLtInline(prev, c1 + c2)) 1L else 0L + w[5] = prev + dCarry = c3b + c4b + + dLo = a3 * a3 + dHi = umulhFused(a3, a3) + s = w[6] + dLo + c1 = if (uLtInline(s, w[6])) 1L else 0L + s += dCarry + c2 = if (uLtInline(s, dCarry)) 1L else 0L + w[6] = s + prev = w[7] + dHi + prev += c1 + c2 + w[7] = prev + + // Reduction + reduceWideInline(out, w) { x, y -> umulhFused(x, y) } +} diff --git a/quartz/src/commonMain/kotlin/com/vitorpamplona/quartz/utils/secp256k1/FieldP.kt b/quartz/src/commonMain/kotlin/com/vitorpamplona/quartz/utils/secp256k1/FieldP.kt index fec7200e5b..af5456b140 100644 --- a/quartz/src/commonMain/kotlin/com/vitorpamplona/quartz/utils/secp256k1/FieldP.kt +++ b/quartz/src/commonMain/kotlin/com/vitorpamplona/quartz/utils/secp256k1/FieldP.kt @@ -68,7 +68,7 @@ internal object FieldP { if (carry != 0) { // Overflow past 2^256: add 2^256 mod p = 2^32 + 977 = 0x1000003D1 val s1 = out[0] + 4294968273L - val c1 = if (uLt(s1, out[0])) 1L else 0L + val c1 = if (uLtInline(s1, out[0])) 1L else 0L out[0] = s1 if (c1 != 0L) { out[1]++ @@ -95,7 +95,7 @@ internal object FieldP { if (borrow != 0) { // Add P = [P0, -1, -1, -1]. val s0 = out[0] + P0 - val c0 = if (uLt(s0, out[0])) 1L else 0L + val c0 = if (uLtInline(s0, out[0])) 1L else 0L out[0] = s0 // For limbs 1-3: adding P[i]=-1 with carry c: // c=1 → result unchanged, carry out=1 (identity propagation) @@ -173,7 +173,7 @@ internal object FieldP { } // P - a: limb 0 is P0 - a[0], limbs 1-3 are (-1) - a[i] = ~a[i] out[0] = P0 - a[0] - val borrow = if (uLt(P0, a[0])) 1L else 0L + val borrow = if (uLtInline(P0, a[0])) 1L else 0L // ~a[i] - borrow. New borrow only if ~a[i] == 0 (i.e., a[i] == -1) and borrow == 1 out[1] = a[1].inv() - borrow val b1 = if (a[1] == -1L && borrow != 0L) 1L else 0L @@ -200,28 +200,28 @@ internal object FieldP { // Conditional add: out = a + (P & mask), unrolled // Limb 0 s1 = a[0] + p0 - c1 = if (uLt(s1, a[0])) 1L else 0L + c1 = if (uLtInline(s1, a[0])) 1L else 0L out[0] = s1 var carry = c1 // Limb 1 s1 = a[1] + mask - c1 = if (uLt(s1, a[1])) 1L else 0L + c1 = if (uLtInline(s1, a[1])) 1L else 0L s2 = s1 + carry - c2 = if (uLt(s2, s1)) 1L else 0L + c2 = if (uLtInline(s2, s1)) 1L else 0L out[1] = s2 carry = c1 + c2 // Limb 2 s1 = a[2] + mask - c1 = if (uLt(s1, a[2])) 1L else 0L + c1 = if (uLtInline(s1, a[2])) 1L else 0L s2 = s1 + carry - c2 = if (uLt(s2, s1)) 1L else 0L + c2 = if (uLtInline(s2, s1)) 1L else 0L out[2] = s2 carry = c1 + c2 // Limb 3 s1 = a[3] + mask - c1 = if (uLt(s1, a[3])) 1L else 0L + c1 = if (uLtInline(s1, a[3])) 1L else 0L s2 = s1 + carry - c2 = if (uLt(s2, s1)) 1L else 0L + c2 = if (uLtInline(s2, s1)) 1L else 0L out[3] = s2 carry = c1 + c2 @@ -405,7 +405,7 @@ internal object FieldP { hcLo = w[4] * c hcHi = unsignedMultiplyHigh(w[4], c) s1 = w[0] + hcLo - c1 = if (uLt(s1, w[0])) 1L else 0L + c1 = if (uLtInline(s1, w[0])) 1L else 0L out[0] = s1 var carry = hcHi + c1 @@ -413,9 +413,9 @@ internal object FieldP { hcLo = w[5] * c hcHi = unsignedMultiplyHigh(w[5], c) s1 = w[1] + hcLo - c1 = if (uLt(s1, w[1])) 1L else 0L + c1 = if (uLtInline(s1, w[1])) 1L else 0L s2 = s1 + carry - c2 = if (uLt(s2, s1)) 1L else 0L + c2 = if (uLtInline(s2, s1)) 1L else 0L out[1] = s2 carry = hcHi + c1 + c2 @@ -423,9 +423,9 @@ internal object FieldP { hcLo = w[6] * c hcHi = unsignedMultiplyHigh(w[6], c) s1 = w[2] + hcLo - c1 = if (uLt(s1, w[2])) 1L else 0L + c1 = if (uLtInline(s1, w[2])) 1L else 0L s2 = s1 + carry - c2 = if (uLt(s2, s1)) 1L else 0L + c2 = if (uLtInline(s2, s1)) 1L else 0L out[2] = s2 carry = hcHi + c1 + c2 @@ -433,9 +433,9 @@ internal object FieldP { hcLo = w[7] * c hcHi = unsignedMultiplyHigh(w[7], c) s1 = w[3] + hcLo - c1 = if (uLt(s1, w[3])) 1L else 0L + c1 = if (uLtInline(s1, w[3])) 1L else 0L s2 = s1 + carry - c2 = if (uLt(s2, s1)) 1L else 0L + c2 = if (uLtInline(s2, s1)) 1L else 0L out[3] = s2 carry = hcHi + c1 + c2 @@ -444,21 +444,21 @@ internal object FieldP { val ccLo = carry * c val ccHi = unsignedMultiplyHigh(carry, c) s1 = out[0] + ccLo - c1 = if (uLt(s1, out[0])) 1L else 0L + c1 = if (uLtInline(s1, out[0])) 1L else 0L out[0] = s1 // Propagate carry (unrolled, with early exit) var prop = ccHi + c1 if (prop != 0L) { s1 = out[1] + prop - prop = if (uLt(s1, out[1])) 1L else 0L + prop = if (uLtInline(s1, out[1])) 1L else 0L out[1] = s1 if (prop != 0L) { s1 = out[2] + prop - prop = if (uLt(s1, out[2])) 1L else 0L + prop = if (uLtInline(s1, out[2])) 1L else 0L out[2] = s1 if (prop != 0L) { s1 = out[3] + prop - prop = if (uLt(s1, out[3])) 1L else 0L + prop = if (uLtInline(s1, out[3])) 1L else 0L out[3] = s1 } } @@ -466,7 +466,7 @@ internal object FieldP { // Overflow past 256 bits: 2^256 ≡ C (mod p) if (prop != 0L) { s1 = out[0] + c - c1 = if (uLt(s1, out[0])) 1L else 0L + c1 = if (uLtInline(s1, out[0])) 1L else 0L out[0] = s1 if (c1 != 0L) { out[1]++ diff --git a/quartz/src/commonMain/kotlin/com/vitorpamplona/quartz/utils/secp256k1/PointTypes.kt b/quartz/src/commonMain/kotlin/com/vitorpamplona/quartz/utils/secp256k1/PointTypes.kt index 6f90ecfe57..40f1e22ec9 100644 --- a/quartz/src/commonMain/kotlin/com/vitorpamplona/quartz/utils/secp256k1/PointTypes.kt +++ b/quartz/src/commonMain/kotlin/com/vitorpamplona/quartz/utils/secp256k1/PointTypes.kt @@ -46,7 +46,7 @@ internal class MutablePoint( @JvmField val y: LongArray = LongArray(4), @JvmField val z: LongArray = LongArray(4), ) { - fun isInfinity(): Boolean = U256.isZero(z) + fun isInfinity(): Boolean = (z[0] or z[1] or z[2] or z[3]) == 0L fun setInfinity() { for (i in 0 until 4) { @@ -58,17 +58,17 @@ internal class MutablePoint( } fun copyFrom(other: MutablePoint) { - other.x.copyInto(x) - other.y.copyInto(y) - other.z.copyInto(z) + other.x.copyInto(x, 0, 0, 4) + other.y.copyInto(y, 0, 0, 4) + other.z.copyInto(z, 0, 0, 4) } fun setAffine( ax: LongArray, ay: LongArray, ) { - ax.copyInto(x) - ay.copyInto(y) + ax.copyInto(x, 0, 0, 4) + ay.copyInto(y, 0, 0, 4) z[0] = 1L for (i in 1 until 4) z[i] = 0L } diff --git a/quartz/src/commonMain/kotlin/com/vitorpamplona/quartz/utils/secp256k1/ScalarN.kt b/quartz/src/commonMain/kotlin/com/vitorpamplona/quartz/utils/secp256k1/ScalarN.kt index 7981f95f1a..cbbdb20008 100644 --- a/quartz/src/commonMain/kotlin/com/vitorpamplona/quartz/utils/secp256k1/ScalarN.kt +++ b/quartz/src/commonMain/kotlin/com/vitorpamplona/quartz/utils/secp256k1/ScalarN.kt @@ -210,9 +210,9 @@ internal object ScalarN { 0L } val s1 = hiTimesNC[i] + loVal - val c1 = if (uLt(s1, hiTimesNC[i])) 1L else 0L + val c1 = if (uLtInline(s1, hiTimesNC[i])) 1L else 0L val s2 = s1 + carry - val c2 = if (uLt(s2, s1)) 1L else 0L + val c2 = if (uLtInline(s2, s1)) 1L else 0L w[i] = s2 carry = c1 + c2 } @@ -250,9 +250,9 @@ internal object ScalarN { saved3 } val s1 = loVal + hi2NC[i] - val c1 = if (uLt(s1, loVal)) 1L else 0L + val c1 = if (uLtInline(s1, loVal)) 1L else 0L val s2 = s1 + c2 - val cc = if (uLt(s2, s1)) 1L else 0L + val cc = if (uLtInline(s2, s1)) 1L else 0L out[i] = s2 c2 = c1 + cc } @@ -264,13 +264,13 @@ internal object ScalarN { val c1lo = ov * N_COMPLEMENT[1] val c1hi = unsignedMultiplyHigh(ov, N_COMPLEMENT[1]) val s0 = out[0] + c0lo - val carry0 = if (uLt(s0, out[0])) 1L else 0L + val carry0 = if (uLtInline(s0, out[0])) 1L else 0L out[0] = s0 val s1 = out[1] + c0hi + c1lo + carry0 - val carry1 = if (uLt(s1, out[1])) 1L else 0L + val carry1 = if (uLtInline(s1, out[1])) 1L else 0L out[1] = s1 val s2 = out[2] + c1hi + ov + carry1 - val carry2 = if (uLt(s2, out[2])) 1L else 0L + val carry2 = if (uLtInline(s2, out[2])) 1L else 0L out[2] = s2 out[3] += carry2 } diff --git a/quartz/src/commonMain/kotlin/com/vitorpamplona/quartz/utils/secp256k1/Secp256k1.kt b/quartz/src/commonMain/kotlin/com/vitorpamplona/quartz/utils/secp256k1/Secp256k1.kt index f61221af51..9eea4728c3 100644 --- a/quartz/src/commonMain/kotlin/com/vitorpamplona/quartz/utils/secp256k1/Secp256k1.kt +++ b/quartz/src/commonMain/kotlin/com/vitorpamplona/quartz/utils/secp256k1/Secp256k1.kt @@ -65,18 +65,9 @@ object Secp256k1 { // The tag prefixes SHA256(tag) || SHA256(tag) are constant per tag string. // We precompute them once to save 2 SHA256 calls per sign/verify operation. - private val CHALLENGE_PREFIX: ByteArray by lazy { - val h = sha256("BIP0340/challenge".encodeToByteArray()) - h + h - } - private val AUX_PREFIX: ByteArray by lazy { - val h = sha256("BIP0340/aux".encodeToByteArray()) - h + h - } - private val NONCE_PREFIX: ByteArray by lazy { - val h = sha256("BIP0340/nonce".encodeToByteArray()) - h + h - } + private val CHALLENGE_PREFIX: ByteArray = sha256("BIP0340/challenge".encodeToByteArray()).let { it + it } + private val AUX_PREFIX: ByteArray = sha256("BIP0340/aux".encodeToByteArray()).let { it + it } + private val NONCE_PREFIX: ByteArray = sha256("BIP0340/nonce".encodeToByteArray()).let { it + it } // ==================== Pubkey decompression cache ==================== // @@ -119,8 +110,8 @@ object Secp256k1 { val cached = pubkeyCache[slot] if (cached != null && cached.keyBytes.contentEquals(pub)) { // Cache hit — copy pre-computed coordinates - cached.px.copyInto(outX) - cached.py.copyInto(outY) + cached.px.copyInto(outX, 0, 0, 4) + cached.py.copyInto(outY, 0, 0, 4) return true } @@ -136,10 +127,11 @@ object Secp256k1 { /** Create a 65-byte uncompressed public key (04 || x || y) from a 32-byte secret key. */ fun pubkeyCreate(seckey: ByteArray): ByteArray { require(seckey.size == 32) - val scalar = U256.fromBytes(seckey) - require(ScalarN.isValid(scalar)) val sc = ECPoint.getScratch() - ECPoint.mulG(sc.entryResult, scalar) + val scalar = sc.scalarTmp1 + U256.fromBytesInto(scalar, seckey, 0) + require(ScalarN.isValid(scalar)) + ECPoint.mulG(sc.entryResult, scalar, sc) check(ECPoint.toAffine(sc.entryResult, sc.entryPx, sc.entryPy, sc)) return KeyCodec.serializeUncompressed(sc.entryPx, sc.entryPy) } @@ -242,13 +234,12 @@ object Secp256k1 { auxrand: ByteArray?, ): ByteArray { require(seckey.size == 32) - // Allocate d0 separately — signSchnorrInternal uses all scalar scratch buffers. + // d0 must survive through mulG (which destroys splitK*) and signSchnorrInternal + // (which uses scalarTmp*). The allocation is negligible vs mulG cost (~100μs). val d0 = U256.fromBytes(seckey) require(ScalarN.isValid(d0)) - - // Derive public key (one G multiplication + one inversion) val sc = ECPoint.getScratch() - ECPoint.mulG(sc.entryResult, d0) + ECPoint.mulG(sc.entryResult, d0, sc) check(ECPoint.toAffine(sc.entryResult, sc.entryPx, sc.entryPy, sc)) // Allocate xOnlyPub — signSchnorrInternal reuses bytesTmp1/2 internally. @@ -279,7 +270,10 @@ object Secp256k1 { auxrand: ByteArray?, ): ByteArray { require(seckey.size == 32 && compressedPub.size == 33) - val d0 = U256.fromBytes(seckey) + // Use zInv for d0 — not used by signSchnorrInternal. + val sc = ECPoint.getScratch() + val d0 = sc.zInv + U256.fromBytesInto(d0, seckey, 0) require(ScalarN.isValid(d0)) val hasEvenY = compressedPub[0] == 0x02.toByte() val xOnlyPub = compressedPub.copyOfRange(1, 33) @@ -302,7 +296,10 @@ object Secp256k1 { auxrand: ByteArray?, ): ByteArray { require(seckey.size == 32 && xOnlyPub.size == 32) - val d0 = U256.fromBytes(seckey) + // Use zInv for d0 — not used by signSchnorrInternal. + val sc = ECPoint.getScratch() + val d0 = sc.zInv + U256.fromBytesInto(d0, seckey, 0) require(ScalarN.isValid(d0)) // BIP-340: x-only pubkeys always have even y return signSchnorrInternal(data, d0, xOnlyPub, true, auxrand) @@ -337,8 +334,8 @@ object Secp256k1 { if (auxrand != null) { require(auxrand.size == 32) // Build AUX_PREFIX + auxrand in scratch hashBuf (avoids concatenation alloc) - AUX_PREFIX.copyInto(sc.hashBuf, 0) - auxrand.copyInto(sc.hashBuf, 64) + AUX_PREFIX.copyInto(sc.hashBuf, 0, 0, 64) + auxrand.copyInto(sc.hashBuf, 64, 0, 32) sha256Into(sc.bytesTmp2, sc.hashBuf, 96) // XOR d with auxHash — reuse limb scratch U256.fromBytesInto(sc.scalarTmp1, dBytes, 0) @@ -353,10 +350,10 @@ object Secp256k1 { // Build nonce input. Reuse hashBuf if it fits. val nonceLen = 64 + 32 + 32 + data.size val nonceInput = if (nonceLen <= sc.hashBuf.size) sc.hashBuf else ByteArray(nonceLen) - NONCE_PREFIX.copyInto(nonceInput, 0) + NONCE_PREFIX.copyInto(nonceInput, 0, 0, 64) tBytes.copyInto(nonceInput, 64, 0, 32) - pBytes.copyInto(nonceInput, 96) - data.copyInto(nonceInput, 128) + pBytes.copyInto(nonceInput, 96, 0, 32) + data.copyInto(nonceInput, 128, 0, data.size) sha256Into(sc.bytesTmp2, nonceInput, nonceLen) // rand → bytesTmp2 U256.fromBytesInto(sc.scalarTmp1, sc.bytesTmp2, 0) ScalarN.reduceTo(sc.scalarTmp1, sc.scalarTmp1) @@ -364,7 +361,7 @@ object Secp256k1 { require(!U256.isZero(k0)) // R = k0·G - ECPoint.mulG(sc.entryResult, k0) + ECPoint.mulG(sc.entryResult, k0, sc) val rx = sc.entryPx val ry = sc.entryPy check(ECPoint.toAffine(sc.entryResult, rx, ry, sc)) @@ -380,10 +377,10 @@ object Secp256k1 { // Challenge: e = H(R || P || msg) — reuse hashBuf val chalLen = 64 + 32 + 32 + data.size val chalInput = if (chalLen <= sc.hashBuf.size) sc.hashBuf else ByteArray(chalLen) - CHALLENGE_PREFIX.copyInto(chalInput, 0) + CHALLENGE_PREFIX.copyInto(chalInput, 0, 0, 64) U256.toBytesInto(rx, chalInput, 64) - pBytes.copyInto(chalInput, 96) - data.copyInto(chalInput, 128) + pBytes.copyInto(chalInput, 96, 0, 32) + data.copyInto(chalInput, 128, 0, data.size) sha256Into(sc.bytesTmp1, chalInput, chalLen) // eHash → bytesTmp1 U256.fromBytesInto(sc.scalarTmp3, sc.bytesTmp1, 0) ScalarN.reduceTo(sc.scalarTmp3, sc.scalarTmp3) @@ -422,45 +419,8 @@ object Secp256k1 { pub: ByteArray, ): Boolean { if (signature.size != 64 || pub.size != 32) return false - - // Use thread-local scratch to avoid per-verify allocations. - // Saves ~10 LongArray(4) + 2 MutablePoint = ~14 object allocations per call. val sc = ECPoint.getScratch() - if (!liftXCached(sc.entryPx, sc.entryPy, pub)) return false - - val r = sc.entryTmp - U256.fromBytesInto(r, signature, 0) - if (U256.cmp(r, FieldP.P) >= 0) return false - val s = sc.entryTmp2 - U256.fromBytesInto(s, signature, 32) - if (U256.cmp(s, ScalarN.N) >= 0) return false - - // Build challenge hash input. Reuse scratch byte buffer if message fits, - // otherwise allocate (rare for Nostr: event IDs are 32 bytes → total 160). - val hashLen = 64 + 32 + 32 + data.size - val hashInput = if (hashLen <= sc.hashBuf.size) sc.hashBuf else ByteArray(hashLen) - CHALLENGE_PREFIX.copyInto(hashInput, 0) - signature.copyInto(hashInput, 64, 0, 32) // r bytes from signature - pub.copyInto(hashInput, 96) - data.copyInto(hashInput, 128) - val eHash = sha256Into(sc.bytesTmp1, hashInput, hashLen) - // Reuse zInv for e (safe: zInv not used until toAffine, which we skip here) - val e = sc.zInv - U256.fromBytesInto(e, eHash, 0) - if (U256.cmp(e, ScalarN.N) >= 0) U256.subTo(e, e, ScalarN.N) // inline reduce - - // Q = s·G + (-e)·P via Shamir's trick - ScalarN.negTo(e, e) // negate in-place - sc.entryPoint.setAffine(sc.entryPx, sc.entryPy) // copies px/py, so entryPx is free - ECPoint.mulDoubleG(sc.entryResult, s, sc.entryPoint, e) - - if (sc.entryResult.isInfinity()) return false - - // Check x-coordinate in Jacobian FIRST (2 field ops, no inversion): X/Z² == r → X == r·Z². - val w = sc.w - FieldP.sqr(sc.zInv2, sc.entryResult.z, w) // Z² - FieldP.mul(sc.zInv3, r, sc.zInv2, w) // r·Z² - if (U256.cmp(sc.entryResult.x, sc.zInv3) != 0) return false // x mismatch → reject fast + if (!verifySchnorrCore(signature, data, pub, sc)) return false // x matches — check y-parity (requires inversion, ~270 field ops) FieldP.inv(sc.zInv, sc.entryResult.z) @@ -497,8 +457,29 @@ object Secp256k1 { pub: ByteArray, ): Boolean { if (signature.size != 64 || pub.size != 32) return false - val sc = ECPoint.getScratch() + return verifySchnorrCore(signature, data, pub, sc) + } + + /** + * Shared core of Schnorr verification: validates inputs, computes + * Q = s·G + (-e)·P via Shamir's trick, and checks that Q.x matches + * the signature's r value in Jacobian coordinates (no inversion). + * + * Leaves the Jacobian result point in [sc].entryResult for callers + * that need additional checks (e.g., y-parity in [verifySchnorr]). + * + * By extracting this into a single method, the JIT compiles one hot path + * for the expensive mulDoubleG instead of two near-identical method bodies. + * All copyInto calls use explicit parameters to avoid the Kotlin + * copyInto$default bridge (bitmask + 3 branches + arraylength per call). + */ + private fun verifySchnorrCore( + signature: ByteArray, + data: ByteArray, + pub: ByteArray, + sc: PointScratch, + ): Boolean { if (!liftXCached(sc.entryPx, sc.entryPy, pub)) return false val r = sc.entryTmp @@ -508,30 +489,31 @@ object Secp256k1 { U256.fromBytesInto(s, signature, 32) if (U256.cmp(s, ScalarN.N) >= 0) return false - // Build challenge hash + // Build challenge hash input. Reuse scratch byte buffer if message fits, + // otherwise allocate (rare for Nostr: event IDs are 32 bytes → total 160). val hashLen = 64 + 32 + 32 + data.size val hashInput = if (hashLen <= sc.hashBuf.size) sc.hashBuf else ByteArray(hashLen) - CHALLENGE_PREFIX.copyInto(hashInput, 0) + CHALLENGE_PREFIX.copyInto(hashInput, 0, 0, 64) signature.copyInto(hashInput, 64, 0, 32) - pub.copyInto(hashInput, 96) - data.copyInto(hashInput, 128) + pub.copyInto(hashInput, 96, 0, 32) + data.copyInto(hashInput, 128, 0, data.size) val eHash = sha256Into(sc.bytesTmp1, hashInput, hashLen) + // Reuse zInv for e (safe: zInv not used until toAffine, which we skip here) val e = sc.zInv U256.fromBytesInto(e, eHash, 0) - if (U256.cmp(e, ScalarN.N) >= 0) U256.subTo(e, e, ScalarN.N) + if (U256.cmp(e, ScalarN.N) >= 0) U256.subTo(e, e, ScalarN.N) // inline reduce - // Q = s·G + (-e)·P - ScalarN.negTo(e, e) - sc.entryPoint.setAffine(sc.entryPx, sc.entryPy) - ECPoint.mulDoubleG(sc.entryResult, s, sc.entryPoint, e) + // Q = s·G + (-e)·P via Shamir's trick + ScalarN.negTo(e, e) // negate in-place + sc.entryPoint.setAffine(sc.entryPx, sc.entryPy) // copies px/py, so entryPx is free + ECPoint.mulDoubleG(sc.entryResult, s, sc.entryPoint, e, sc) if (sc.entryResult.isInfinity()) return false - // Jacobian x-check only — no inversion, no y-parity check. - // Saves ~270 field ops (~14% of verify). + // Jacobian x-check: X == r·Z² (2 field ops, no inversion) val w = sc.w - FieldP.sqr(sc.zInv2, sc.entryResult.z, w) - FieldP.mul(sc.zInv3, r, sc.zInv2, w) + FieldP.sqr(sc.zInv2, sc.entryResult.z, w) // Z² + FieldP.mul(sc.zInv3, r, sc.zInv2, w) // r·Z² return U256.cmp(sc.entryResult.x, sc.zInv3) == 0 } @@ -543,9 +525,18 @@ object Secp256k1 { tweak: ByteArray, ): ByteArray { require(seckey.size == 32 && tweak.size == 32) - val result = ScalarN.add(U256.fromBytes(seckey), U256.fromBytes(tweak)) - require(!U256.isZero(result) && U256.cmp(result, ScalarN.N) < 0) - return U256.toBytes(result) + // Use thread-local scratch to avoid 2 intermediate LongArray(4) allocations. + // Old path: fromBytes (alloc) + fromBytes (alloc) + add (alloc) + toBytes (alloc) = 4 allocs + // New path: fromBytesInto (scratch) + fromBytesInto (scratch) + addTo (scratch) + toBytes = 1 alloc + val sc = ECPoint.getScratch() + val a = sc.entryTmp + val b = sc.entryTmp2 + val r = sc.scalarTmp1 + U256.fromBytesInto(a, seckey, 0) + U256.fromBytesInto(b, tweak, 0) + ScalarN.addTo(r, a, b) + require(!U256.isZero(r) && U256.cmp(r, ScalarN.N) < 0) + return U256.toBytes(r) } /** Multiply a public key by a scalar. Used for ECDH shared secret derivation. */ @@ -556,11 +547,12 @@ object Secp256k1 { require(tweak.size == 32) val sc = ECPoint.getScratch() check(KeyCodec.parsePublicKey(pubkey, sc.entryPx, sc.entryPy)) - val scalar = U256.fromBytes(tweak) + val scalar = sc.scalarTmp1 + U256.fromBytesInto(scalar, tweak, 0) require(ScalarN.isValid(scalar)) sc.entryPoint.setAffine(sc.entryPx, sc.entryPy) - ECPoint.mul(sc.entryResult, sc.entryPoint, scalar) + ECPoint.mul(sc.entryResult, sc.entryPoint, scalar, sc) check(ECPoint.toAffine(sc.entryResult, sc.entryPx, sc.entryPy, sc)) return if (pubkey.size == 33) { @@ -591,7 +583,8 @@ object Secp256k1 { val sc = ECPoint.getScratch() U256.fromBytesInto(sc.entryTmp, xOnlyPub, 0) require(U256.cmp(sc.entryTmp, FieldP.P) < 0) - val k = U256.fromBytes(scalar) + val k = sc.scalarTmp1 + U256.fromBytesInto(k, scalar, 0) require(ScalarN.isValid(k)) // Compute y = sqrt(x³ + 7). We need SOME valid y for EC point operations, @@ -601,7 +594,7 @@ object Secp256k1 { } sc.entryPoint.setAffine(sc.entryPx, sc.entryPy) - ECPoint.mul(sc.entryResult, sc.entryPoint, k) + ECPoint.mul(sc.entryResult, sc.entryPoint, k, sc) check(ECPoint.toAffineX(sc.entryResult, sc.entryPx, sc)) return U256.toBytes(sc.entryPx) } @@ -696,10 +689,10 @@ object Secp256k1 { // Compute challenge eᵢ = H(rᵢ || pub || msgᵢ) using scratch buffers val hashLen = 64 + 32 + 32 + msg.size val hashInput = if (hashLen <= sc.hashBuf.size) sc.hashBuf else ByteArray(hashLen) - CHALLENGE_PREFIX.copyInto(hashInput, 0) + CHALLENGE_PREFIX.copyInto(hashInput, 0, 0, 64) sig.copyInto(hashInput, 64, 0, 32) - pub.copyInto(hashInput, 96) - msg.copyInto(hashInput, 128) + pub.copyInto(hashInput, 96, 0, 32) + msg.copyInto(hashInput, 128, 0, msg.size) sha256Into(sc.bytesTmp1, hashInput, hashLen) U256.fromBytesInto(e, sc.bytesTmp1, 0) ScalarN.reduceTo(e, e) diff --git a/quartz/src/commonMain/kotlin/com/vitorpamplona/quartz/utils/secp256k1/U256.kt b/quartz/src/commonMain/kotlin/com/vitorpamplona/quartz/utils/secp256k1/U256.kt index b3e6dacd1f..8c5486dac0 100644 --- a/quartz/src/commonMain/kotlin/com/vitorpamplona/quartz/utils/secp256k1/U256.kt +++ b/quartz/src/commonMain/kotlin/com/vitorpamplona/quartz/utils/secp256k1/U256.kt @@ -106,7 +106,7 @@ internal object U256 { ): Int { for (i in 3 downTo 0) { if (a[i] != b[i]) { - return if (uLt(a[i], b[i])) -1 else 1 + return if (uLtInline(a[i], b[i])) -1 else 1 } } return 0 @@ -125,31 +125,31 @@ internal object U256 { // Limb 0 (no carry input) s1 = a[0] + b[0] - c1 = if (uLt(s1, a[0])) 1L else 0L + c1 = if (uLtInline(s1, a[0])) 1L else 0L out[0] = s1 var carry = c1 // Limb 1 s1 = a[1] + b[1] - c1 = if (uLt(s1, a[1])) 1L else 0L + c1 = if (uLtInline(s1, a[1])) 1L else 0L s2 = s1 + carry - c2 = if (uLt(s2, s1)) 1L else 0L + c2 = if (uLtInline(s2, s1)) 1L else 0L out[1] = s2 carry = c1 + c2 // Limb 2 s1 = a[2] + b[2] - c1 = if (uLt(s1, a[2])) 1L else 0L + c1 = if (uLtInline(s1, a[2])) 1L else 0L s2 = s1 + carry - c2 = if (uLt(s2, s1)) 1L else 0L + c2 = if (uLtInline(s2, s1)) 1L else 0L out[2] = s2 carry = c1 + c2 // Limb 3 s1 = a[3] + b[3] - c1 = if (uLt(s1, a[3])) 1L else 0L + c1 = if (uLtInline(s1, a[3])) 1L else 0L s2 = s1 + carry - c2 = if (uLt(s2, s1)) 1L else 0L + c2 = if (uLtInline(s2, s1)) 1L else 0L out[3] = s2 carry = c1 + c2 @@ -169,31 +169,31 @@ internal object U256 { // Limb 0 (no borrow input) d1 = a[0] - b[0] - c1 = if (uLt(a[0], b[0])) 1L else 0L + c1 = if (uLtInline(a[0], b[0])) 1L else 0L out[0] = d1 var borrow = c1 // Limb 1 d1 = a[1] - b[1] - c1 = if (uLt(a[1], b[1])) 1L else 0L + c1 = if (uLtInline(a[1], b[1])) 1L else 0L d2 = d1 - borrow - c2 = if (uLt(d1, borrow)) 1L else 0L + c2 = if (uLtInline(d1, borrow)) 1L else 0L out[1] = d2 borrow = c1 + c2 // Limb 2 d1 = a[2] - b[2] - c1 = if (uLt(a[2], b[2])) 1L else 0L + c1 = if (uLtInline(a[2], b[2])) 1L else 0L d2 = d1 - borrow - c2 = if (uLt(d1, borrow)) 1L else 0L + c2 = if (uLtInline(d1, borrow)) 1L else 0L out[2] = d2 borrow = c1 + c2 // Limb 3 d1 = a[3] - b[3] - c1 = if (uLt(a[3], b[3])) 1L else 0L + c1 = if (uLtInline(a[3], b[3])) 1L else 0L d2 = d1 - borrow - c2 = if (uLt(d1, borrow)) 1L else 0L + c2 = if (uLtInline(d1, borrow)) 1L else 0L out[3] = d2 borrow = c1 + c2 @@ -235,19 +235,19 @@ internal object U256 { lo = a0 * b1 s = lo + carry - c1 = if (uLt(s, lo)) 1L else 0L + c1 = if (uLtInline(s, lo)) 1L else 0L out[1] = s carry = unsignedMultiplyHigh(a0, b1) + c1 lo = a0 * b2 s = lo + carry - c1 = if (uLt(s, lo)) 1L else 0L + c1 = if (uLtInline(s, lo)) 1L else 0L out[2] = s carry = unsignedMultiplyHigh(a0, b2) + c1 lo = a0 * b3 s = lo + carry - c1 = if (uLt(s, lo)) 1L else 0L + c1 = if (uLtInline(s, lo)) 1L else 0L out[3] = s out[4] = unsignedMultiplyHigh(a0, b3) + c1 @@ -256,7 +256,7 @@ internal object U256 { hi = unsignedMultiplyHigh(a1, b0) prev = out[1] s = prev + lo - c1 = if (uLt(s, prev)) 1L else 0L + c1 = if (uLtInline(s, prev)) 1L else 0L out[1] = s carry = hi + c1 @@ -264,9 +264,9 @@ internal object U256 { hi = unsignedMultiplyHigh(a1, b1) prev = out[2] s = prev + lo - c1 = if (uLt(s, prev)) 1L else 0L + c1 = if (uLtInline(s, prev)) 1L else 0L s += carry - c2 = if (uLt(s, carry)) 1L else 0L + c2 = if (uLtInline(s, carry)) 1L else 0L out[2] = s carry = hi + c1 + c2 @@ -274,9 +274,9 @@ internal object U256 { hi = unsignedMultiplyHigh(a1, b2) prev = out[3] s = prev + lo - c1 = if (uLt(s, prev)) 1L else 0L + c1 = if (uLtInline(s, prev)) 1L else 0L s += carry - c2 = if (uLt(s, carry)) 1L else 0L + c2 = if (uLtInline(s, carry)) 1L else 0L out[3] = s carry = hi + c1 + c2 @@ -284,9 +284,9 @@ internal object U256 { hi = unsignedMultiplyHigh(a1, b3) prev = out[4] s = prev + lo - c1 = if (uLt(s, prev)) 1L else 0L + c1 = if (uLtInline(s, prev)) 1L else 0L s += carry - c2 = if (uLt(s, carry)) 1L else 0L + c2 = if (uLtInline(s, carry)) 1L else 0L out[4] = s out[5] = hi + c1 + c2 @@ -295,7 +295,7 @@ internal object U256 { hi = unsignedMultiplyHigh(a2, b0) prev = out[2] s = prev + lo - c1 = if (uLt(s, prev)) 1L else 0L + c1 = if (uLtInline(s, prev)) 1L else 0L out[2] = s carry = hi + c1 @@ -303,9 +303,9 @@ internal object U256 { hi = unsignedMultiplyHigh(a2, b1) prev = out[3] s = prev + lo - c1 = if (uLt(s, prev)) 1L else 0L + c1 = if (uLtInline(s, prev)) 1L else 0L s += carry - c2 = if (uLt(s, carry)) 1L else 0L + c2 = if (uLtInline(s, carry)) 1L else 0L out[3] = s carry = hi + c1 + c2 @@ -313,9 +313,9 @@ internal object U256 { hi = unsignedMultiplyHigh(a2, b2) prev = out[4] s = prev + lo - c1 = if (uLt(s, prev)) 1L else 0L + c1 = if (uLtInline(s, prev)) 1L else 0L s += carry - c2 = if (uLt(s, carry)) 1L else 0L + c2 = if (uLtInline(s, carry)) 1L else 0L out[4] = s carry = hi + c1 + c2 @@ -323,9 +323,9 @@ internal object U256 { hi = unsignedMultiplyHigh(a2, b3) prev = out[5] s = prev + lo - c1 = if (uLt(s, prev)) 1L else 0L + c1 = if (uLtInline(s, prev)) 1L else 0L s += carry - c2 = if (uLt(s, carry)) 1L else 0L + c2 = if (uLtInline(s, carry)) 1L else 0L out[5] = s out[6] = hi + c1 + c2 @@ -334,7 +334,7 @@ internal object U256 { hi = unsignedMultiplyHigh(a3, b0) prev = out[3] s = prev + lo - c1 = if (uLt(s, prev)) 1L else 0L + c1 = if (uLtInline(s, prev)) 1L else 0L out[3] = s carry = hi + c1 @@ -342,9 +342,9 @@ internal object U256 { hi = unsignedMultiplyHigh(a3, b1) prev = out[4] s = prev + lo - c1 = if (uLt(s, prev)) 1L else 0L + c1 = if (uLtInline(s, prev)) 1L else 0L s += carry - c2 = if (uLt(s, carry)) 1L else 0L + c2 = if (uLtInline(s, carry)) 1L else 0L out[4] = s carry = hi + c1 + c2 @@ -352,9 +352,9 @@ internal object U256 { hi = unsignedMultiplyHigh(a3, b2) prev = out[5] s = prev + lo - c1 = if (uLt(s, prev)) 1L else 0L + c1 = if (uLtInline(s, prev)) 1L else 0L s += carry - c2 = if (uLt(s, carry)) 1L else 0L + c2 = if (uLtInline(s, carry)) 1L else 0L out[5] = s carry = hi + c1 + c2 @@ -362,9 +362,9 @@ internal object U256 { hi = unsignedMultiplyHigh(a3, b3) prev = out[6] s = prev + lo - c1 = if (uLt(s, prev)) 1L else 0L + c1 = if (uLtInline(s, prev)) 1L else 0L s += carry - c2 = if (uLt(s, carry)) 1L else 0L + c2 = if (uLtInline(s, carry)) 1L else 0L out[6] = s out[7] = hi + c1 + c2 } @@ -401,13 +401,13 @@ internal object U256 { lo = a0 * a2 s = lo + carry - c1 = if (uLt(s, lo)) 1L else 0L + c1 = if (uLtInline(s, lo)) 1L else 0L out[2] = s carry = unsignedMultiplyHigh(a0, a2) + c1 lo = a0 * a3 s = lo + carry - c1 = if (uLt(s, lo)) 1L else 0L + c1 = if (uLtInline(s, lo)) 1L else 0L out[3] = s out[4] = unsignedMultiplyHigh(a0, a3) + c1 @@ -416,7 +416,7 @@ internal object U256 { hi = unsignedMultiplyHigh(a1, a2) prev = out[3] s = prev + lo - c1 = if (uLt(s, prev)) 1L else 0L + c1 = if (uLtInline(s, prev)) 1L else 0L out[3] = s carry = hi + c1 @@ -424,9 +424,9 @@ internal object U256 { hi = unsignedMultiplyHigh(a1, a3) prev = out[4] s = prev + lo - c1 = if (uLt(s, prev)) 1L else 0L + c1 = if (uLtInline(s, prev)) 1L else 0L s += carry - c2 = if (uLt(s, carry)) 1L else 0L + c2 = if (uLtInline(s, carry)) 1L else 0L out[4] = s out[5] = hi + c1 + c2 @@ -435,7 +435,7 @@ internal object U256 { hi = unsignedMultiplyHigh(a2, a3) prev = out[5] s = prev + lo - c1 = if (uLt(s, prev)) 1L else 0L + c1 = if (uLtInline(s, prev)) 1L else 0L out[5] = s out[6] = hi + c1 @@ -467,7 +467,7 @@ internal object U256 { hi = unsignedMultiplyHigh(a0, a0) out[0] = lo // out[0] was 0 s = out[1] + hi - c1 = if (uLt(s, out[1])) 1L else 0L + c1 = if (uLtInline(s, out[1])) 1L else 0L out[1] = s var dCarry = c1 @@ -475,14 +475,14 @@ internal object U256 { lo = a1 * a1 hi = unsignedMultiplyHigh(a1, a1) s = out[2] + lo - c1 = if (uLt(s, out[2])) 1L else 0L + c1 = if (uLtInline(s, out[2])) 1L else 0L s += dCarry - c2 = if (uLt(s, dCarry)) 1L else 0L + c2 = if (uLtInline(s, dCarry)) 1L else 0L out[2] = s prev = out[3] + hi - val c3a = if (uLt(prev, out[3])) 1L else 0L + val c3a = if (uLtInline(prev, out[3])) 1L else 0L prev += c1 + c2 - val c4a = if (uLt(prev, c1 + c2)) 1L else 0L + val c4a = if (uLtInline(prev, c1 + c2)) 1L else 0L out[3] = prev dCarry = c3a + c4a @@ -490,14 +490,14 @@ internal object U256 { lo = a2 * a2 hi = unsignedMultiplyHigh(a2, a2) s = out[4] + lo - c1 = if (uLt(s, out[4])) 1L else 0L + c1 = if (uLtInline(s, out[4])) 1L else 0L s += dCarry - c2 = if (uLt(s, dCarry)) 1L else 0L + c2 = if (uLtInline(s, dCarry)) 1L else 0L out[4] = s prev = out[5] + hi - val c3b = if (uLt(prev, out[5])) 1L else 0L + val c3b = if (uLtInline(prev, out[5])) 1L else 0L prev += c1 + c2 - val c4b = if (uLt(prev, c1 + c2)) 1L else 0L + val c4b = if (uLtInline(prev, c1 + c2)) 1L else 0L out[5] = prev dCarry = c3b + c4b @@ -505,12 +505,12 @@ internal object U256 { lo = a3 * a3 hi = unsignedMultiplyHigh(a3, a3) s = out[6] + lo - c1 = if (uLt(s, out[6])) 1L else 0L + c1 = if (uLtInline(s, out[6])) 1L else 0L s += dCarry - c2 = if (uLt(s, dCarry)) 1L else 0L + c2 = if (uLtInline(s, dCarry)) 1L else 0L out[6] = s prev = out[7] + hi - val c3c = if (uLt(prev, out[7])) 1L else 0L + val c3c = if (uLtInline(prev, out[7])) 1L else 0L prev += c1 + c2 out[7] = prev } @@ -602,6 +602,6 @@ internal object U256 { out: LongArray, a: LongArray, ) { - a.copyInto(out) + a.copyInto(out, 0, 0, 4) } } diff --git a/quartz/src/nativeMain/kotlin/com/vitorpamplona/quartz/utils/secp256k1/FieldMulPlatform.native.kt b/quartz/src/nativeMain/kotlin/com/vitorpamplona/quartz/utils/secp256k1/FieldMulPlatform.native.kt index 84946c8324..720688720a 100644 --- a/quartz/src/nativeMain/kotlin/com/vitorpamplona/quartz/utils/secp256k1/FieldMulPlatform.native.kt +++ b/quartz/src/nativeMain/kotlin/com/vitorpamplona/quartz/utils/secp256k1/FieldMulPlatform.native.kt @@ -30,7 +30,7 @@ internal actual fun fieldMulReduce( b: LongArray, w: LongArray, ) { - fieldMulReduceWith(out, a, b, w) { x, y -> unsignedMultiplyHighFallback(x, y) } + fieldMulReduceFused(out, a, b, w) } internal actual fun fieldSqrReduce( @@ -38,5 +38,5 @@ internal actual fun fieldSqrReduce( a: LongArray, w: LongArray, ) { - fieldSqrReduceWith(out, a, w) { x, y -> unsignedMultiplyHighFallback(x, y) } + fieldSqrReduceFused(out, a, w) }