diff --git a/quartz/src/commonMain/kotlin/com/vitorpamplona/quartz/utils/secp256k1/Glv.kt b/quartz/src/commonMain/kotlin/com/vitorpamplona/quartz/utils/secp256k1/Glv.kt index f22292831..64a76585b 100644 --- a/quartz/src/commonMain/kotlin/com/vitorpamplona/quartz/utils/secp256k1/Glv.kt +++ b/quartz/src/commonMain/kotlin/com/vitorpamplona/quartz/utils/secp256k1/Glv.kt @@ -39,10 +39,13 @@ package com.vitorpamplona.quartz.utils.secp256k1 internal object Glv { /** β: cube root of unity mod p. φ(x,y) = (β·x, y). */ - val BETA = longArrayOf( - -4523465429756870162L, -7138124642204153451L, - 7954561588662645993L, 8856726876819556112L, - ) + val BETA = + longArrayOf( + -4523465429756870162L, + -7138124642204153451L, + 7954561588662645993L, + 8856726876819556112L, + ) // ==================== GLV Scalar Decomposition ==================== @@ -63,7 +66,8 @@ internal object Glv { return Split( if (neg1) ScalarN.neg(r1) else r1, if (neg2) ScalarN.neg(r2) else r2, - neg1, neg2, + neg1, + neg2, ) } @@ -75,7 +79,11 @@ internal object Glv { * * The working copy is extended to handle carries past maxBits. */ - fun wnaf(scalar: LongArray, w: Int, maxBits: Int): IntArray { + fun wnaf( + scalar: LongArray, + w: Int, + maxBits: Int, + ): IntArray { val totalBits = maxBits + w val sLimbs = maxOf((totalBits + 63) / 64, scalar.size) val result = IntArray(totalBits) @@ -101,7 +109,10 @@ internal object Glv { // ==================== Internal Helpers ==================== /** Multiply two 256-bit numbers, return result >> 384 (rounded). */ - private fun mulShift384(k: LongArray, g: LongArray): LongArray { + private fun mulShift384( + k: LongArray, + g: LongArray, + ): LongArray { val wide = LongArray(8) U256.mulWide(wide, k, g) val result = LongArray(4) @@ -115,7 +126,11 @@ internal object Glv { return result } - private fun getBitsVar(s: LongArray, bitPos: Int, count: Int): Int { + private fun getBitsVar( + s: LongArray, + bitPos: Int, + count: Int, + ): Int { if (count == 0) return 0 val limb = bitPos / 64 val shift = bitPos % 64 @@ -126,7 +141,10 @@ internal object Glv { return (r and ((1L shl count) - 1L)).toInt() } - private fun addBitTo(s: LongArray, bitPos: Int) { + private fun addBitTo( + s: LongArray, + bitPos: Int, + ) { val limb = bitPos / 64 if (limb >= s.size) return val addVal = 1L shl (bitPos % 64) @@ -140,26 +158,46 @@ internal object Glv { // ==================== Constants (from libsecp256k1) ==================== - private val MINUS_LAMBDA = longArrayOf( - -2247357714951666737L, -6304834983940376126L, - 6546514211138018212L, -6008836872998760673L, - ) - private val G1 = longArrayOf( - -1687969588364726223L, 4443515802769476223L, - -1698823648040391915L, 3496713202691238861L, - ) - private val G2 = longArrayOf( - 1545214808910233457L, 2455034284347819718L, - 8022177200260244676L, -1998614352016537560L, - ) - private val MINUS_B1 = longArrayOf( - 8022177200260244675L, -1998614352016537560L, 0L, 0L, - ) - private val MINUS_B2 = longArrayOf( - -2925706260434037204L, -8491525256057179027L, -2L, -1L, - ) - private val N_HALF = longArrayOf( - -2312264954237214560L, 6725966010171805725L, - -1L, 9223372036854775807L, - ) + private val MINUS_LAMBDA = + longArrayOf( + -2247357714951666737L, + -6304834983940376126L, + 6546514211138018212L, + -6008836872998760673L, + ) + private val G1 = + longArrayOf( + -1687969588364726223L, + 4443515802769476223L, + -1698823648040391915L, + 3496713202691238861L, + ) + private val G2 = + longArrayOf( + 1545214808910233457L, + 2455034284347819718L, + 8022177200260244676L, + -1998614352016537560L, + ) + private val MINUS_B1 = + longArrayOf( + 8022177200260244675L, + -1998614352016537560L, + 0L, + 0L, + ) + private val MINUS_B2 = + longArrayOf( + -2925706260434037204L, + -8491525256057179027L, + -2L, + -1L, + ) + private val N_HALF = + longArrayOf( + -2312264954237214560L, + 6725966010171805725L, + -1L, + 9223372036854775807L, + ) } diff --git a/quartz/src/commonMain/kotlin/com/vitorpamplona/quartz/utils/secp256k1/Point.kt b/quartz/src/commonMain/kotlin/com/vitorpamplona/quartz/utils/secp256k1/Point.kt index 10aa61072..f03152476 100644 --- a/quartz/src/commonMain/kotlin/com/vitorpamplona/quartz/utils/secp256k1/Point.kt +++ b/quartz/src/commonMain/kotlin/com/vitorpamplona/quartz/utils/secp256k1/Point.kt @@ -115,14 +115,20 @@ internal class AffinePoint( internal object ECPoint { // ==================== Generator point G ==================== - val GX = longArrayOf( - 6481385041966929816L, 188021827762530521L, - 6170039885052185351L, 8772561819708210092L, - ) - val GY = longArrayOf( - -7185545363635252040L, -209500633525038055L, - 6747795201694173352L, 5204712524664259685L, - ) + val GX = + longArrayOf( + 6481385041966929816L, + 188021827762530521L, + 6170039885052185351L, + 8772561819708210092L, + ) + val GY = + longArrayOf( + -7185545363635252040L, + -209500633525038055L, + 6747795201694173352L, + 5204712524664259685L, + ) /** Curve constant b = 7 in y² = x³ + 7. */ private val B = longArrayOf(7L, 0L, 0L, 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 37143f924..d2f44b1ed 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 @@ -118,21 +118,30 @@ internal object ScalarN { result[i] = s2 c2 = c1 + cc } - // Handle remaining overflow - var overflow = c2 - for (i in 4 until 8) { - overflow += hi2NC[i].toULong().toLong() // approximate - } - if (overflow != 0L) { - // overflow × N_COMPLEMENT is small, fold it in - val corr0 = overflow * N_COMPLEMENT[0] - val s = result[0] + corr0 - val c = if (s.toULong() < result[0].toULong()) 1L else 0L - result[0] = s - if (c != 0L || overflow * N_COMPLEMENT[1] != 0L) { - val s1 = result[1] + overflow * N_COMPLEMENT[1] + c - result[1] = s1 - } + // Handle remaining overflow from hi2NC[4..7] + carry + // hi2NC[4..7] should be small (hi2 is ~129 bits, NC is ~129 bits → product ≤ 258 bits) + // So hi2NC[4] might be non-zero but hi2NC[5..7] should be zero. + // Fold: overflow * N_COMPLEMENT into result + var ov = c2 + hi2NC[4] + for (i in 5 until 8) ov += hi2NC[i] + if (ov != 0L) { + // ov × NC[0] + val c0lo = ov * N_COMPLEMENT[0] + val c0hi = unsignedMultiplyHigh(ov, N_COMPLEMENT[0]) + // ov × NC[1] + val c1lo = ov * N_COMPLEMENT[1] + val c1hi = unsignedMultiplyHigh(ov, N_COMPLEMENT[1]) + // ov × NC[2] = ov × 1 = ov + val s0 = result[0] + c0lo + val carry0 = if (s0.toULong() < result[0].toULong()) 1L else 0L + result[0] = s0 + val s1 = result[1] + c0hi + c1lo + carry0 + val carry1 = if (s1.toULong() < result[1].toULong()) 1L else 0L + result[1] = s1 + val s2 = result[2] + c1hi + ov + carry1 + val carry2 = if (s2.toULong() < result[2].toULong()) 1L else 0L + result[2] = s2 + result[3] += carry2 } while (U256.cmp(result, N) >= 0) U256.subTo(result, result, N) diff --git a/quartz/src/commonTest/kotlin/com/vitorpamplona/quartz/utils/secp256k1/GlvTest.kt b/quartz/src/commonTest/kotlin/com/vitorpamplona/quartz/utils/secp256k1/GlvTest.kt index 0b4c2fc86..557128d3a 100644 --- a/quartz/src/commonTest/kotlin/com/vitorpamplona/quartz/utils/secp256k1/GlvTest.kt +++ b/quartz/src/commonTest/kotlin/com/vitorpamplona/quartz/utils/secp256k1/GlvTest.kt @@ -70,7 +70,7 @@ class GlvTest { val k = hex("67e56582298859ddae725f972992a07c6c4fb9f62a8fff58ce3ca926a1063530") val split = Glv.splitScalar(k) // Upper 4 limbs should be zero for a proper 128-bit half-scalar - for (i in 4 until 8) { + for (i in 2 until 4) { assertEquals(0, split.k1[i], "k1 limb $i should be 0") assertEquals(0, split.k2[i], "k2 limb $i should be 0") } @@ -137,7 +137,7 @@ class GlvTest { val k = hex("e907831f80848d1069a5371b402410364bdf1c5f8307b0084c55f1ce2dca8215") val digits = Glv.wnaf(k, 5, 256) val reconstructed = reconstructWnaf(digits) - for (i in 0 until 8) assertEquals(k[i], reconstructed[i], "Limb $i mismatch") + for (i in 0 until 4) assertEquals(k[i], reconstructed[i], "Limb $i mismatch") } @Test @@ -147,7 +147,7 @@ class GlvTest { val k = hex("944946c56e0d133f326db0b0645544a04bfcc5a0f447ad6d0227958414d8ba73") val digits = Glv.wnaf(k, 5, 256) val reconstructed = reconstructWnaf(digits) - for (i in 0 until 8) assertEquals(k[i], reconstructed[i], "Limb $i mismatch for carry test") + for (i in 0 until 4) assertEquals(k[i], reconstructed[i], "Limb $i mismatch for carry test") } @Test @@ -227,12 +227,22 @@ class GlvTest { val s = acc[0] + d.toLong() val c = if (s.toULong() < acc[0].toULong()) 1L else 0L acc[0] = s - if (c != 0L) for (j in 1 until 4) { acc[j]++; if (acc[j] != 0L) break } + if (c != 0L) { + for (j in 1 until 4) { + acc[j]++ + if (acc[j] != 0L) break + } + } } else if (d < 0) { val s = acc[0] - (-d).toLong() val b = if (acc[0].toULong() < (-d).toULong()) 1L else 0L acc[0] = s - if (b != 0L) for (j in 1 until 4) { acc[j]--; if (acc[j] != -1L) break } + if (b != 0L) { + for (j in 1 until 4) { + acc[j]-- + if (acc[j] != -1L) break + } + } } } return acc