fix: reduceWide round-2 carry overflow — fixes NIP-44 conversation key for edge-case scalars

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
This commit is contained in:
Claude
2026-04-06 13:12:18 +00:00
parent 18b8df2d2d
commit 36a7ae147e
3 changed files with 291 additions and 97 deletions
@@ -26,18 +26,23 @@ package com.vitorpamplona.quartz.utils.secp256k1
*/ */
internal object FieldP { internal object FieldP {
// p = 0xFFFFFFFFFFFFFFFFFFFFFFFFFFFFFFFFFFFFFFFFFFFFFFFFFFFFFFFEFFFFFC2F // p = 0xFFFFFFFFFFFFFFFFFFFFFFFFFFFFFFFFFFFFFFFFFFFFFFFFFFFFFFFEFFFFFC2F
val P = longArrayOf( val P =
-4294968273L, // 0xFFFFFFFEFFFFFC2F longArrayOf(
-1L, // 0xFFFFFFFFFFFFFFFF -4294968273L, // 0xFFFFFFFEFFFFFC2F
-1L, // 0xFFFFFFFFFFFFFFFF -1L, // 0xFFFFFFFFFFFFFFFF
-1L, // 0xFFFFFFFFFFFFFFFF -1L, // 0xFFFFFFFFFFFFFFFF
) -1L, // 0xFFFFFFFFFFFFFFFF
)
private val wide = ThreadLocal.withInitial { LongArray(8) } private val wide = ThreadLocal.withInitial { LongArray(8) }
// ==================== Core arithmetic ==================== // ==================== 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) val carry = U256.addTo(out, a, b)
if (carry != 0) { if (carry != 0) {
// Overflow past 2^256: add 2^256 mod p = 2^32 + 977 = 0x1000003D1 // Overflow past 2^256: add 2^256 mod p = 2^32 + 977 = 0x1000003D1
@@ -55,24 +60,38 @@ internal object FieldP {
reduceSelf(out) 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) val borrow = U256.subTo(out, a, b)
if (borrow != 0) U256.addTo(out, out, P) 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() val w = wide.get()
U256.mulWide(w, a, b) U256.mulWide(w, a, b)
reduceWide(out, w) reduceWide(out, w)
} }
fun sqr(out: LongArray, a: LongArray) { fun sqr(
out: LongArray,
a: LongArray,
) {
val w = wide.get() val w = wide.get()
U256.sqrWide(w, a) U256.sqrWide(w, a)
reduceWide(out, w) reduceWide(out, w)
} }
fun neg(out: LongArray, a: LongArray) { fun neg(
out: LongArray,
a: LongArray,
) {
if (U256.isZero(a)) { if (U256.isZero(a)) {
for (i in 0 until 4) out[i] = 0L for (i in 0 until 4) out[i] = 0L
} else { } 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). * 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 val mask = -(a[0] and 1L) // all 1s if odd, all 0s if even
var carry = 0L var carry = 0L
for (i in 0 until 4) { for (i in 0 until 4) {
@@ -104,60 +126,114 @@ internal object FieldP {
// ==================== Inversion and square root (optimized addition chains) ==================== // ==================== Inversion and square root (optimized addition chains) ====================
fun inv(out: LongArray, a: LongArray) { fun inv(
out: LongArray,
a: LongArray,
) {
require(!U256.isZero(a)) require(!U256.isZero(a))
val x2 = LongArray(4); val x3 = LongArray(4); val x6 = LongArray(4) val x2 = LongArray(4)
val x9 = LongArray(4); val x11 = LongArray(4); val x22 = LongArray(4) val x3 = LongArray(4)
val x44 = LongArray(4); val x88 = LongArray(4); val x176 = LongArray(4) val x6 = LongArray(4)
val x220 = LongArray(4); val x223 = 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(x2, a)
sqr(x3, x2); mul(x3, x3, a) mul(x2, x2, a)
sqrN(x6, x3, 3); mul(x6, x6, x3) sqr(x3, x2)
sqrN(x9, x6, 3); mul(x9, x9, x3) mul(x3, x3, a)
sqrN(x11, x9, 2); mul(x11, x11, x2) sqrN(x6, x3, 3)
sqrN(x22, x11, 11); mul(x22, x22, x11) mul(x6, x6, x3)
sqrN(x44, x22, 22); mul(x44, x44, x22) sqrN(x9, x6, 3)
sqrN(x88, x44, 44); mul(x88, x88, x44) mul(x9, x9, x3)
sqrN(x176, x88, 88); mul(x176, x176, x88) sqrN(x11, x9, 2)
sqrN(x220, x176, 44); mul(x220, x220, x44) mul(x11, x11, x2)
sqrN(x223, x220, 3); mul(x223, x223, x3) 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, x223, 23)
sqrN(out, out, 5); mul(out, out, a) mul(out, out, x22)
sqrN(out, out, 3); mul(out, out, x2) sqrN(out, out, 5)
sqrN(out, out, 2); mul(out, out, a) 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 { fun sqrt(
val x2 = LongArray(4); val x3 = LongArray(4); val x6 = LongArray(4) out: LongArray,
val x9 = LongArray(4); val x11 = LongArray(4); val x22 = LongArray(4) a: LongArray,
val x44 = LongArray(4); val x88 = LongArray(4); val x176 = LongArray(4) ): Boolean {
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(x2, a)
sqr(x3, x2); mul(x3, x3, a) mul(x2, x2, a)
sqrN(x6, x3, 3); mul(x6, x6, x3) sqr(x3, x2)
sqrN(x9, x6, 3); mul(x9, x9, x3) mul(x3, x3, a)
sqrN(x11, x9, 2); mul(x11, x11, x2) sqrN(x6, x3, 3)
sqrN(x22, x11, 11); mul(x22, x22, x11) mul(x6, x6, x3)
sqrN(x44, x22, 22); mul(x44, x44, x22) sqrN(x9, x6, 3)
sqrN(x88, x44, 44); mul(x88, x88, x44) mul(x9, x9, x3)
sqrN(x176, x88, 88); mul(x176, x176, x88) sqrN(x11, x9, 2)
sqrN(x220, x176, 44); mul(x220, x220, x44) mul(x11, x11, x2)
sqrN(x223, x220, 3); mul(x223, x223, x3) 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, x223, 23)
sqrN(out, out, 6); mul(out, out, x2) mul(out, out, x22)
sqrN(out, out, 6)
mul(out, out, x2)
sqrN(out, out, 2) sqrN(out, out, 2)
val check = LongArray(4) val check = LongArray(4)
mul(check, out, out) 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 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) U256.copyInto(out, a)
repeat(n) { sqr(out, out) } 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. * 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 * 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. * 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 // Round 1: acc = lo + hi × C
val c = 4294968273L // 2^32 + 977 val c = 4294968273L // 2^32 + 977
var carry = 0L var carry = 0L
for (i in 0 until 4) { for (i in 0 until 4) {
val hiC_lo = w[i + 4] * c val hcLo = w[i + 4] * c
val hiC_hi = unsignedMultiplyHigh(w[i + 4], c) val hcHi = unsignedMultiplyHigh(w[i + 4], c)
// acc = w[i] + hiC_lo + carry // acc = w[i] + hcLo + carry
val s1 = w[i] + hiC_lo val s1 = w[i] + hcLo
val c1 = if (s1.toULong() < w[i].toULong()) 1L else 0L val c1 = if (s1.toULong() < w[i].toULong()) 1L else 0L
val s2 = s1 + carry val s2 = s1 + carry
val c2 = if (s2.toULong() < s1.toULong()) 1L else 0L val c2 = if (s2.toULong() < s1.toULong()) 1L else 0L
out[i] = s2 out[i] = s2
carry = hiC_hi + c1 + c2 carry = hcHi + c1 + c2
} }
// Round 2: if carry > 0, fold carry × C back in // Round 2: if carry > 0, fold carry × C back in
if (carry != 0L) { if (carry != 0L) {
val cc_lo = carry * c val ccLo = carry * c
val cc_hi = unsignedMultiplyHigh(carry, c) val ccHi = unsignedMultiplyHigh(carry, c)
val s1 = out[0] + cc_lo val s1 = out[0] + ccLo
val c1 = if (s1.toULong() < out[0].toULong()) 1L else 0L val c1 = if (s1.toULong() < out[0].toULong()) 1L else 0L
out[0] = s1 out[0] = s1
var prop = cc_hi + c1 var prop = ccHi + c1
for (i in 1 until 4) { for (i in 1 until 4) {
if (prop == 0L) break if (prop == 0L) break
val s = out[i] + prop val s = out[i] + prop
prop = if (s.toULong() < out[i].toULong()) 1L else 0L prop = if (s.toULong() < out[i].toULong()) 1L else 0L
out[i] = s 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 // Final: at most one subtraction of p
@@ -214,12 +313,59 @@ internal object FieldP {
// ==================== Convenience wrappers ==================== // ==================== Convenience wrappers ====================
fun add(a: LongArray, b: LongArray): LongArray { val r = LongArray(4); add(r, a, b); return r } fun add(
fun sub(a: LongArray, b: LongArray): LongArray { val r = LongArray(4); sub(r, a, b); return r } a: LongArray,
fun mul(a: LongArray, b: LongArray): LongArray { val r = LongArray(4); mul(r, a, b); return r } b: LongArray,
fun sqr(a: LongArray): LongArray { val r = LongArray(4); sqr(r, a); return r } ): LongArray {
fun neg(a: LongArray): LongArray { val r = LongArray(4); neg(r, a); return r } val r = LongArray(4)
fun inv(a: LongArray): LongArray { val r = LongArray(4); inv(r, a); return r } add(r, a, b)
fun sqrt(a: LongArray): LongArray? { val r = LongArray(4); return if (sqrt(r, a)) r else null } return r
fun reduce(a: LongArray): LongArray { val r = a.copyOf(); reduceSelf(r); 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
}
} }
@@ -24,26 +24,45 @@ package com.vitorpamplona.quartz.utils.secp256k1
* Arithmetic modulo the secp256k1 group order n using LongArray(4) limbs. * Arithmetic modulo the secp256k1 group order n using LongArray(4) limbs.
*/ */
internal object ScalarN { internal object ScalarN {
val N = longArrayOf( val N =
-4624529908474429119L, -4994812053365940165L, -2L, -1L, longArrayOf(
) -4624529908474429119L,
-4994812053365940165L,
-2L,
-1L,
)
private val N_COMPLEMENT = longArrayOf( private val N_COMPLEMENT =
4624529908474429119L, 4994812053365940164L, 1L, 0L, longArrayOf(
) 4624529908474429119L,
4994812053365940164L,
1L,
0L,
)
private val N_MINUS_2 = longArrayOf( private val N_MINUS_2 =
-4624529908474429121L, -4994812053365940165L, -2L, -1L, longArrayOf(
) -4624529908474429121L,
-4994812053365940165L,
-2L,
-1L,
)
fun isValid(a: LongArray): Boolean = !U256.isZero(a) && U256.cmp(a, N) < 0 fun isValid(a: LongArray): Boolean = !U256.isZero(a) && U256.cmp(a, N) < 0
fun reduce(a: LongArray): LongArray = fun reduce(a: LongArray): LongArray =
if (U256.cmp(a, N) >= 0) { if (U256.cmp(a, N) >= 0) {
val r = LongArray(4); U256.subTo(r, a, N); r val r = LongArray(4)
} else a 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 r = LongArray(4)
val carry = U256.addTo(r, a, b) val carry = U256.addTo(r, a, b)
if (carry != 0) U256.addTo(r, r, N_COMPLEMENT) if (carry != 0) U256.addTo(r, r, N_COMPLEMENT)
@@ -51,14 +70,20 @@ internal object ScalarN {
return r return r
} }
fun sub(a: LongArray, b: LongArray): LongArray { fun sub(
a: LongArray,
b: LongArray,
): LongArray {
val r = LongArray(4) val r = LongArray(4)
val borrow = U256.subTo(r, a, b) val borrow = U256.subTo(r, a, b)
if (borrow != 0) U256.addTo(r, r, N) if (borrow != 0) U256.addTo(r, r, N)
return r return r
} }
fun mul(a: LongArray, b: LongArray): LongArray { fun mul(
a: LongArray,
b: LongArray,
): LongArray {
val w = LongArray(8) val w = LongArray(8)
U256.mulWide(w, a, b) U256.mulWide(w, a, b)
return reduceWide(w) return reduceWide(w)
@@ -66,7 +91,9 @@ internal object ScalarN {
fun neg(a: LongArray): LongArray { fun neg(a: LongArray): LongArray {
if (U256.isZero(a)) return LongArray(4) 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 { 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. * Uses hi × 2^256 ≡ hi × N_COMPLEMENT (mod n). N_COMPLEMENT is ~129 bits.
*/ */
private fun reduceWide(w: LongArray): LongArray { private fun reduceWide(w: LongArray): LongArray {
val lo = LongArray(4); val hi = LongArray(4) val lo = LongArray(4)
for (i in 0 until 4) { lo[i] = w[i]; hi[i] = w[i + 4] } val hi = LongArray(4)
if (U256.isZero(hi)) { reduceSelf(lo); return lo } 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 // Round 1: lo + hi × N_COMPLEMENT
val hiTimesNC = LongArray(8) val hiTimesNC = LongArray(8)
@@ -102,9 +136,16 @@ internal object ScalarN {
} }
// Round 2 if still > 256 bits // Round 2 if still > 256 bits
val lo2 = LongArray(4); val hi2 = LongArray(4) val lo2 = LongArray(4)
for (i in 0 until 4) { lo2[i] = sum[i]; hi2[i] = sum[i + 4] } val hi2 = LongArray(4)
if (U256.isZero(hi2)) { reduceSelf(lo2); return lo2 } 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) val hi2NC = LongArray(8)
U256.mulWide(hi2NC, hi2, N_COMPLEMENT) U256.mulWide(hi2NC, hi2, N_COMPLEMENT)
@@ -148,12 +189,18 @@ internal object ScalarN {
return result return result
} }
private fun powModN(base: LongArray, exp: LongArray): LongArray { private fun powModN(
base: LongArray,
exp: LongArray,
): LongArray {
val result = LongArray(4) val result = LongArray(4)
val b = base.copyOf() val b = base.copyOf()
var highBit = 255 var highBit = 255
while (highBit >= 0 && !U256.testBit(exp, highBit)) highBit-- 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) U256.copyInto(result, b)
for (i in highBit - 1 downTo 0) { for (i in highBit - 1 downTo 0) {
val sq = mul(result, result) val sq = mul(result, result)
@@ -37,10 +37,11 @@ import com.vitorpamplona.quartz.utils.sha256.sha256
* - [pubKeyTweakMul]: ECDH shared secrets (NIP-04, NIP-44) * - [pubKeyTweakMul]: ECDH shared secrets (NIP-04, NIP-44)
* *
* Performance on JVM (vs native C/JNI secp256k1): * 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). * 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 * Uses 4×64-bit limbs (LongArray(4)) with Math.multiplyHigh for 64×64→128-bit products
* limbs with 64 inner products per field multiply, vs C's 5×52-bit with 25). * (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, * All algorithmic optimizations from libsecp256k1 are implemented: GLV endomorphism,
* wNAF encoding, Shamir's trick, comb method, and optimized addition chains. * 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). * 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 { fun secKeyVerify(seckey: ByteArray): Boolean {
if (seckey.size != 32) return false if (seckey.size != 32) return false