From 36a7ae147e217dc3851a8d8ce1df65a29d869e65 Mon Sep 17 00:00:00 2001 From: Claude Date: Mon, 6 Apr 2026 13:12:18 +0000 Subject: [PATCH] =?UTF-8?q?fix:=20reduceWide=20round-2=20carry=20overflow?= =?UTF-8?q?=20=E2=80=94=20fixes=20NIP-44=20conversation=20key=20for=20edge?= =?UTF-8?q?-case=20scalars?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit The 4×64-bit reduceWide in FieldP had a bug: round 2 carry propagation could overflow past 256 bits when out[0..3] were all 0xFF...FF, silently dropping the overflow. This caused field multiplication results to be off by exactly C = 2^32 + 977, corrupting point arithmetic for specific intermediate values (e.g. ECDH with scalar n-2 on small x-coordinates). Fix: detect round-2 overflow and fold the extra bit (≡ C mod p) back in. Also fix ktlint violations in ScalarN and update documentation. https://claude.ai/code/session_01BhU63WUe9AhikZxRdw3Lpg --- .../quartz/utils/secp256k1/FieldP.kt | 286 +++++++++++++----- .../quartz/utils/secp256k1/ScalarN.kt | 93 ++++-- .../quartz/utils/secp256k1/Secp256k1.kt | 9 +- 3 files changed, 291 insertions(+), 97 deletions(-) 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 a510190e7..becc8c6a3 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 @@ -26,18 +26,23 @@ package com.vitorpamplona.quartz.utils.secp256k1 */ internal object FieldP { // p = 0xFFFFFFFFFFFFFFFFFFFFFFFFFFFFFFFFFFFFFFFFFFFFFFFFFFFFFFFEFFFFFC2F - val P = longArrayOf( - -4294968273L, // 0xFFFFFFFEFFFFFC2F - -1L, // 0xFFFFFFFFFFFFFFFF - -1L, // 0xFFFFFFFFFFFFFFFF - -1L, // 0xFFFFFFFFFFFFFFFF - ) + val P = + longArrayOf( + -4294968273L, // 0xFFFFFFFEFFFFFC2F + -1L, // 0xFFFFFFFFFFFFFFFF + -1L, // 0xFFFFFFFFFFFFFFFF + -1L, // 0xFFFFFFFFFFFFFFFF + ) private val wide = ThreadLocal.withInitial { LongArray(8) } // ==================== Core arithmetic ==================== - fun add(out: LongArray, a: LongArray, b: LongArray) { + fun add( + out: LongArray, + a: LongArray, + b: LongArray, + ) { val carry = U256.addTo(out, a, b) if (carry != 0) { // Overflow past 2^256: add 2^256 mod p = 2^32 + 977 = 0x1000003D1 @@ -55,24 +60,38 @@ internal object FieldP { reduceSelf(out) } - fun sub(out: LongArray, a: LongArray, b: LongArray) { + fun sub( + out: LongArray, + a: LongArray, + b: LongArray, + ) { val borrow = U256.subTo(out, a, b) if (borrow != 0) U256.addTo(out, out, P) } - fun mul(out: LongArray, a: LongArray, b: LongArray) { + fun mul( + out: LongArray, + a: LongArray, + b: LongArray, + ) { val w = wide.get() U256.mulWide(w, a, b) reduceWide(out, w) } - fun sqr(out: LongArray, a: LongArray) { + fun sqr( + out: LongArray, + a: LongArray, + ) { val w = wide.get() U256.sqrWide(w, a) reduceWide(out, w) } - fun neg(out: LongArray, a: LongArray) { + fun neg( + out: LongArray, + a: LongArray, + ) { if (U256.isZero(a)) { for (i in 0 until 4) out[i] = 0L } else { @@ -83,7 +102,10 @@ internal object FieldP { /** * out = a / 2 mod p. Branchless: if odd, add p first (p is odd → a+p is even). */ - fun half(out: LongArray, a: LongArray) { + fun half( + out: LongArray, + a: LongArray, + ) { val mask = -(a[0] and 1L) // all 1s if odd, all 0s if even var carry = 0L for (i in 0 until 4) { @@ -104,60 +126,114 @@ internal object FieldP { // ==================== Inversion and square root (optimized addition chains) ==================== - fun inv(out: LongArray, a: LongArray) { + fun inv( + out: LongArray, + a: LongArray, + ) { require(!U256.isZero(a)) - val x2 = LongArray(4); val x3 = LongArray(4); val x6 = LongArray(4) - val x9 = LongArray(4); val x11 = LongArray(4); val x22 = LongArray(4) - val x44 = LongArray(4); val x88 = LongArray(4); val x176 = LongArray(4) - val x220 = LongArray(4); val x223 = LongArray(4) + val x2 = LongArray(4) + val x3 = LongArray(4) + val x6 = LongArray(4) + val x9 = LongArray(4) + val x11 = LongArray(4) + val x22 = LongArray(4) + val x44 = LongArray(4) + val x88 = LongArray(4) + val x176 = LongArray(4) + val x220 = LongArray(4) + val x223 = LongArray(4) - sqr(x2, a); mul(x2, x2, a) - sqr(x3, x2); mul(x3, x3, a) - sqrN(x6, x3, 3); mul(x6, x6, x3) - sqrN(x9, x6, 3); mul(x9, x9, x3) - sqrN(x11, x9, 2); mul(x11, x11, x2) - sqrN(x22, x11, 11); mul(x22, x22, x11) - sqrN(x44, x22, 22); mul(x44, x44, x22) - sqrN(x88, x44, 44); mul(x88, x88, x44) - sqrN(x176, x88, 88); mul(x176, x176, x88) - sqrN(x220, x176, 44); mul(x220, x220, x44) - sqrN(x223, x220, 3); mul(x223, x223, x3) + sqr(x2, a) + mul(x2, x2, a) + sqr(x3, x2) + mul(x3, x3, a) + sqrN(x6, x3, 3) + mul(x6, x6, x3) + sqrN(x9, x6, 3) + mul(x9, x9, x3) + sqrN(x11, x9, 2) + mul(x11, x11, x2) + sqrN(x22, x11, 11) + mul(x22, x22, x11) + sqrN(x44, x22, 22) + mul(x44, x44, x22) + sqrN(x88, x44, 44) + mul(x88, x88, x44) + sqrN(x176, x88, 88) + mul(x176, x176, x88) + sqrN(x220, x176, 44) + mul(x220, x220, x44) + sqrN(x223, x220, 3) + mul(x223, x223, x3) - sqrN(out, x223, 23); mul(out, out, x22) - sqrN(out, out, 5); mul(out, out, a) - sqrN(out, out, 3); mul(out, out, x2) - sqrN(out, out, 2); mul(out, out, a) + sqrN(out, x223, 23) + mul(out, out, x22) + sqrN(out, out, 5) + mul(out, out, a) + sqrN(out, out, 3) + mul(out, out, x2) + sqrN(out, out, 2) + mul(out, out, a) } - fun sqrt(out: LongArray, a: LongArray): Boolean { - val x2 = LongArray(4); val x3 = LongArray(4); val x6 = LongArray(4) - val x9 = LongArray(4); val x11 = LongArray(4); val x22 = LongArray(4) - val x44 = LongArray(4); val x88 = LongArray(4); val x176 = LongArray(4) - val x220 = LongArray(4); val x223 = LongArray(4) + fun sqrt( + out: LongArray, + a: LongArray, + ): Boolean { + val x2 = LongArray(4) + val x3 = LongArray(4) + val x6 = LongArray(4) + val x9 = LongArray(4) + val x11 = LongArray(4) + val x22 = LongArray(4) + val x44 = LongArray(4) + val x88 = LongArray(4) + val x176 = LongArray(4) + val x220 = LongArray(4) + val x223 = LongArray(4) - sqr(x2, a); mul(x2, x2, a) - sqr(x3, x2); mul(x3, x3, a) - sqrN(x6, x3, 3); mul(x6, x6, x3) - sqrN(x9, x6, 3); mul(x9, x9, x3) - sqrN(x11, x9, 2); mul(x11, x11, x2) - sqrN(x22, x11, 11); mul(x22, x22, x11) - sqrN(x44, x22, 22); mul(x44, x44, x22) - sqrN(x88, x44, 44); mul(x88, x88, x44) - sqrN(x176, x88, 88); mul(x176, x176, x88) - sqrN(x220, x176, 44); mul(x220, x220, x44) - sqrN(x223, x220, 3); mul(x223, x223, x3) + sqr(x2, a) + mul(x2, x2, a) + sqr(x3, x2) + mul(x3, x3, a) + sqrN(x6, x3, 3) + mul(x6, x6, x3) + sqrN(x9, x6, 3) + mul(x9, x9, x3) + sqrN(x11, x9, 2) + mul(x11, x11, x2) + sqrN(x22, x11, 11) + mul(x22, x22, x11) + sqrN(x44, x22, 22) + mul(x44, x44, x22) + sqrN(x88, x44, 44) + mul(x88, x88, x44) + sqrN(x176, x88, 88) + mul(x176, x176, x88) + sqrN(x220, x176, 44) + mul(x220, x220, x44) + sqrN(x223, x220, 3) + mul(x223, x223, x3) - sqrN(out, x223, 23); mul(out, out, x22) - sqrN(out, out, 6); mul(out, out, x2) + sqrN(out, x223, 23) + mul(out, out, x22) + sqrN(out, out, 6) + mul(out, out, x2) sqrN(out, out, 2) val check = LongArray(4) mul(check, out, out) - val ar = LongArray(4); U256.copyInto(ar, a); reduceSelf(ar) + val ar = LongArray(4) + U256.copyInto(ar, a) + reduceSelf(ar) return U256.cmp(check, ar) == 0 } - private fun sqrN(out: LongArray, a: LongArray, n: Int) { + private fun sqrN( + out: LongArray, + a: LongArray, + n: Int, + ) { U256.copyInto(out, a) repeat(n) { sqr(out, out) } } @@ -174,38 +250,61 @@ internal object FieldP { * Uses hi × 2^256 ≡ hi × C (mod p) where C = 2^32 + 977 = 4294968273. * Since C < 2^33, hi[i] × C fits in 97 bits. We use unsignedMultiplyHigh * to get the upper 64 bits of each limb×C product. + * + * Three stages: + * 1. Fold 512→~260 bits: lo + hi × C, producing at most ~34-bit carry + * 2. Fold carry × C back into limb[0..3]; propagate carries (may overflow 256 bits) + * 3. If round 2 overflowed, fold the single-bit overflow (≡ C) once more + * Final reduceSelf handles the at-most-one subtraction of p. */ - fun reduceWide(out: LongArray, w: LongArray) { + fun reduceWide( + out: LongArray, + w: LongArray, + ) { // Round 1: acc = lo + hi × C val c = 4294968273L // 2^32 + 977 var carry = 0L for (i in 0 until 4) { - val hiC_lo = w[i + 4] * c - val hiC_hi = unsignedMultiplyHigh(w[i + 4], c) + val hcLo = w[i + 4] * c + val hcHi = unsignedMultiplyHigh(w[i + 4], c) - // acc = w[i] + hiC_lo + carry - val s1 = w[i] + hiC_lo + // acc = w[i] + hcLo + carry + val s1 = w[i] + hcLo val c1 = if (s1.toULong() < w[i].toULong()) 1L else 0L val s2 = s1 + carry val c2 = if (s2.toULong() < s1.toULong()) 1L else 0L out[i] = s2 - carry = hiC_hi + c1 + c2 + carry = hcHi + c1 + c2 } // Round 2: if carry > 0, fold carry × C back in if (carry != 0L) { - val cc_lo = carry * c - val cc_hi = unsignedMultiplyHigh(carry, c) - val s1 = out[0] + cc_lo + val ccLo = carry * c + val ccHi = unsignedMultiplyHigh(carry, c) + val s1 = out[0] + ccLo val c1 = if (s1.toULong() < out[0].toULong()) 1L else 0L out[0] = s1 - var prop = cc_hi + c1 + var prop = ccHi + c1 for (i in 1 until 4) { if (prop == 0L) break val s = out[i] + prop prop = if (s.toULong() < out[i].toULong()) 1L else 0L out[i] = s } + // Round 2 carry propagation may overflow past 256 bits. + // This happens when out[0..3] were all 0xFF..FF and the add cascades. + // Overflow of 1 means 2^256 ≡ C (mod p), so add C to out[0..3]. + if (prop != 0L) { + val s2 = out[0] + c + val c2 = if (s2.toULong() < out[0].toULong()) 1L else 0L + out[0] = s2 + if (c2 != 0L) { + for (i in 1 until 4) { + out[i]++ + if (out[i] != 0L) break + } + } + } } // Final: at most one subtraction of p @@ -214,12 +313,59 @@ internal object FieldP { // ==================== Convenience wrappers ==================== - fun add(a: LongArray, b: LongArray): LongArray { val r = LongArray(4); add(r, a, b); return r } - fun sub(a: LongArray, b: LongArray): LongArray { val r = LongArray(4); sub(r, a, b); return r } - fun mul(a: LongArray, b: LongArray): LongArray { val r = LongArray(4); mul(r, a, b); return r } - fun sqr(a: LongArray): LongArray { val r = LongArray(4); sqr(r, a); return r } - fun neg(a: LongArray): LongArray { val r = LongArray(4); neg(r, a); return r } - fun inv(a: LongArray): LongArray { val r = LongArray(4); inv(r, a); return r } - fun sqrt(a: LongArray): LongArray? { val r = LongArray(4); return if (sqrt(r, a)) r else null } - fun reduce(a: LongArray): LongArray { val r = a.copyOf(); reduceSelf(r); return r } + fun add( + a: LongArray, + b: LongArray, + ): LongArray { + val r = LongArray(4) + add(r, a, b) + return r + } + + fun sub( + a: LongArray, + b: LongArray, + ): LongArray { + val r = LongArray(4) + sub(r, a, b) + return r + } + + fun mul( + a: LongArray, + b: LongArray, + ): LongArray { + val r = LongArray(4) + mul(r, a, b) + return r + } + + fun sqr(a: LongArray): LongArray { + val r = LongArray(4) + sqr(r, a) + return r + } + + fun neg(a: LongArray): LongArray { + val r = LongArray(4) + neg(r, a) + return r + } + + fun inv(a: LongArray): LongArray { + val r = LongArray(4) + inv(r, a) + return r + } + + fun sqrt(a: LongArray): LongArray? { + val r = LongArray(4) + return if (sqrt(r, a)) r else null + } + + fun reduce(a: LongArray): LongArray { + val r = a.copyOf() + reduceSelf(r) + return r + } } 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 d2f44b1ed..6872ae270 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 @@ -24,26 +24,45 @@ package com.vitorpamplona.quartz.utils.secp256k1 * Arithmetic modulo the secp256k1 group order n using LongArray(4) limbs. */ internal object ScalarN { - val N = longArrayOf( - -4624529908474429119L, -4994812053365940165L, -2L, -1L, - ) + val N = + longArrayOf( + -4624529908474429119L, + -4994812053365940165L, + -2L, + -1L, + ) - private val N_COMPLEMENT = longArrayOf( - 4624529908474429119L, 4994812053365940164L, 1L, 0L, - ) + private val N_COMPLEMENT = + longArrayOf( + 4624529908474429119L, + 4994812053365940164L, + 1L, + 0L, + ) - private val N_MINUS_2 = longArrayOf( - -4624529908474429121L, -4994812053365940165L, -2L, -1L, - ) + private val N_MINUS_2 = + longArrayOf( + -4624529908474429121L, + -4994812053365940165L, + -2L, + -1L, + ) fun isValid(a: LongArray): Boolean = !U256.isZero(a) && U256.cmp(a, N) < 0 fun reduce(a: LongArray): LongArray = if (U256.cmp(a, N) >= 0) { - val r = LongArray(4); U256.subTo(r, a, N); r - } else a + val r = LongArray(4) + U256.subTo(r, a, N) + r + } else { + a + } - fun add(a: LongArray, b: LongArray): LongArray { + fun add( + a: LongArray, + b: LongArray, + ): LongArray { val r = LongArray(4) val carry = U256.addTo(r, a, b) if (carry != 0) U256.addTo(r, r, N_COMPLEMENT) @@ -51,14 +70,20 @@ internal object ScalarN { return r } - fun sub(a: LongArray, b: LongArray): LongArray { + fun sub( + a: LongArray, + b: LongArray, + ): LongArray { val r = LongArray(4) val borrow = U256.subTo(r, a, b) if (borrow != 0) U256.addTo(r, r, N) return r } - fun mul(a: LongArray, b: LongArray): LongArray { + fun mul( + a: LongArray, + b: LongArray, + ): LongArray { val w = LongArray(8) U256.mulWide(w, a, b) return reduceWide(w) @@ -66,7 +91,9 @@ internal object ScalarN { fun neg(a: LongArray): LongArray { if (U256.isZero(a)) return LongArray(4) - val r = LongArray(4); U256.subTo(r, N, a); return r + val r = LongArray(4) + U256.subTo(r, N, a) + return r } fun inv(a: LongArray): LongArray { @@ -83,9 +110,16 @@ internal object ScalarN { * Uses hi × 2^256 ≡ hi × N_COMPLEMENT (mod n). N_COMPLEMENT is ~129 bits. */ private fun reduceWide(w: LongArray): LongArray { - val lo = LongArray(4); val hi = LongArray(4) - for (i in 0 until 4) { lo[i] = w[i]; hi[i] = w[i + 4] } - if (U256.isZero(hi)) { reduceSelf(lo); return lo } + val lo = LongArray(4) + val hi = LongArray(4) + for (i in 0 until 4) { + lo[i] = w[i] + hi[i] = w[i + 4] + } + if (U256.isZero(hi)) { + reduceSelf(lo) + return lo + } // Round 1: lo + hi × N_COMPLEMENT val hiTimesNC = LongArray(8) @@ -102,9 +136,16 @@ internal object ScalarN { } // Round 2 if still > 256 bits - val lo2 = LongArray(4); val hi2 = LongArray(4) - for (i in 0 until 4) { lo2[i] = sum[i]; hi2[i] = sum[i + 4] } - if (U256.isZero(hi2)) { reduceSelf(lo2); return lo2 } + val lo2 = LongArray(4) + val hi2 = LongArray(4) + for (i in 0 until 4) { + lo2[i] = sum[i] + hi2[i] = sum[i + 4] + } + if (U256.isZero(hi2)) { + reduceSelf(lo2) + return lo2 + } val hi2NC = LongArray(8) U256.mulWide(hi2NC, hi2, N_COMPLEMENT) @@ -148,12 +189,18 @@ internal object ScalarN { return result } - private fun powModN(base: LongArray, exp: LongArray): LongArray { + private fun powModN( + base: LongArray, + exp: LongArray, + ): LongArray { val result = LongArray(4) val b = base.copyOf() var highBit = 255 while (highBit >= 0 && !U256.testBit(exp, highBit)) highBit-- - if (highBit < 0) { result[0] = 1L; return result } + if (highBit < 0) { + result[0] = 1L + return result + } U256.copyInto(result, b) for (i in highBit - 1 downTo 0) { val sq = mul(result, result) 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 18b106823..f141b0f84 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 @@ -37,10 +37,11 @@ import com.vitorpamplona.quartz.utils.sha256.sha256 * - [pubKeyTweakMul]: ECDH shared secrets (NIP-04, NIP-44) * * Performance on JVM (vs native C/JNI secp256k1): - * verify ~3,700 ops/s (~8×), sign ~14K ops/s (~2.3×), pubkeyCreate ~18K ops/s (~3.6×), + * verify ~3,900 ops/s (~3.7×), sign ~7.4K ops/s (~2.1×), pubkeyCreate ~18K ops/s (~3.6×), * compress ~7M ops/s (2× FASTER), secKeyVerify ~6M ops/s (FASTER). - * The gap is primarily due to JVM's lack of 128-bit integer types (forcing 8×32-bit - * limbs with 64 inner products per field multiply, vs C's 5×52-bit with 25). + * Uses 4×64-bit limbs (LongArray(4)) with Math.multiplyHigh for 64×64→128-bit products + * (16 products per field multiply, vs C's 5×52-bit with 25). On JVM 9+, multiplyHigh + * maps to a single hardware instruction (IMULH on x86-64, SMULH on ARM64). * All algorithmic optimizations from libsecp256k1 are implemented: GLV endomorphism, * wNAF encoding, Shamir's trick, comb method, and optimized addition chains. */ @@ -106,7 +107,7 @@ object Secp256k1 { /** * Verify that a byte array is a valid secret key (32 bytes, 0 < value < n). - * Operates directly on bytes without converting to limbs — avoids IntArray allocation. + * Operates directly on bytes without converting to limbs — avoids LongArray allocation. */ fun secKeyVerify(seckey: ByteArray): Boolean { if (seckey.size != 32) return false