perf: eliminate ThreadLocal.get() from hot paths — 27-52% faster across all EC ops

FieldP.mul/sqr were calling ThreadLocal.get() for every invocation (~500+
times per scalar multiplication, ~20-30ns each on JVM). Point operations
(doublePoint, addMixed, addPoints) each did an additional ThreadLocal.get()
for their scratch buffers.

Fix: add overloads that accept a pre-fetched wide buffer (LongArray(8))
and PointScratch. Top-level entry points (mulG, mul, mulDoubleG) fetch
the ThreadLocal once and thread it through all inner calls.

Results (ops/s, vs native JNI):
- pubkeyCreate:       19,163 → 29,205 (+52%, 3.0x → 2.2x)
- signSchnorr cached: 13,007 → 18,397 (+41%, 2.1x → 1.5x)
- signSchnorr:         5,365 →  7,490 (+40%, 5.7x → 3.7x)
- verifySchnorr:       3,840 →  4,873 (+27%, 7.2x → 5.4x)
- ECDH:                5,569 →  7,870 (+41%, 5.5x → 3.8x)

https://claude.ai/code/session_01BhU63WUe9AhikZxRdw3Lpg
This commit is contained in:
Claude
2026-04-06 13:38:01 +00:00
parent 36a7ae147e
commit 547be89577
2 changed files with 208 additions and 130 deletions
@@ -22,7 +22,10 @@ package com.vitorpamplona.quartz.utils.secp256k1
/** /**
* Arithmetic modulo the secp256k1 field prime: p = 2^256 - 2^32 - 977. * Arithmetic modulo the secp256k1 field prime: p = 2^256 - 2^32 - 977.
* Uses LongArray(4) limbs (4×64-bit). Thread-local LongArray(8) scratch for mul/sqr. * Uses LongArray(4) limbs (4×64-bit).
*
* Hot-path mul/sqr accept a pre-fetched LongArray(8) wide buffer to avoid
* ThreadLocal.get() overhead (~20-30ns per call, 500+ calls per scalar mul).
*/ */
internal object FieldP { internal object FieldP {
// p = 0xFFFFFFFFFFFFFFFFFFFFFFFFFFFFFFFFFFFFFFFFFFFFFFFFFFFFFFFEFFFFFC2F // p = 0xFFFFFFFFFFFFFFFFFFFFFFFFFFFFFFFFFFFFFFFFFFFFFFFFFFFFFFFEFFFFFC2F
@@ -36,6 +39,9 @@ internal object FieldP {
private val wide = ThreadLocal.withInitial { LongArray(8) } private val wide = ThreadLocal.withInitial { LongArray(8) }
/** Get a thread-local wide buffer. Call once at the top-level entry point, then pass through. */
fun getWide(): LongArray = wide.get()
// ==================== Core arithmetic ==================== // ==================== Core arithmetic ====================
fun add( fun add(
@@ -69,6 +75,7 @@ internal object FieldP {
if (borrow != 0) U256.addTo(out, out, P) if (borrow != 0) U256.addTo(out, out, P)
} }
/** Multiply with ThreadLocal wide buffer (convenience for non-hot paths). */
fun mul( fun mul(
out: LongArray, out: LongArray,
a: LongArray, a: LongArray,
@@ -79,6 +86,18 @@ internal object FieldP {
reduceWide(out, w) reduceWide(out, w)
} }
/** Multiply with caller-provided wide buffer (hot path — no ThreadLocal lookup). */
fun mul(
out: LongArray,
a: LongArray,
b: LongArray,
w: LongArray,
) {
U256.mulWide(w, a, b)
reduceWide(out, w)
}
/** Square with ThreadLocal wide buffer (convenience for non-hot paths). */
fun sqr( fun sqr(
out: LongArray, out: LongArray,
a: LongArray, a: LongArray,
@@ -88,6 +107,16 @@ internal object FieldP {
reduceWide(out, w) reduceWide(out, w)
} }
/** Square with caller-provided wide buffer (hot path — no ThreadLocal lookup). */
fun sqr(
out: LongArray,
a: LongArray,
w: LongArray,
) {
U256.sqrWide(w, a)
reduceWide(out, w)
}
fun neg( fun neg(
out: LongArray, out: LongArray,
a: LongArray, a: LongArray,
@@ -131,6 +160,7 @@ internal object FieldP {
a: LongArray, a: LongArray,
) { ) {
require(!U256.isZero(a)) require(!U256.isZero(a))
val w = wide.get()
val x2 = LongArray(4) val x2 = LongArray(4)
val x3 = LongArray(4) val x3 = LongArray(4)
val x6 = LongArray(4) val x6 = LongArray(4)
@@ -143,43 +173,44 @@ internal object FieldP {
val x220 = LongArray(4) val x220 = LongArray(4)
val x223 = LongArray(4) val x223 = LongArray(4)
sqr(x2, a) sqr(x2, a, w)
mul(x2, x2, a) mul(x2, x2, a, w)
sqr(x3, x2) sqr(x3, x2, w)
mul(x3, x3, a) mul(x3, x3, a, w)
sqrN(x6, x3, 3) sqrN(x6, x3, 3, w)
mul(x6, x6, x3) mul(x6, x6, x3, w)
sqrN(x9, x6, 3) sqrN(x9, x6, 3, w)
mul(x9, x9, x3) mul(x9, x9, x3, w)
sqrN(x11, x9, 2) sqrN(x11, x9, 2, w)
mul(x11, x11, x2) mul(x11, x11, x2, w)
sqrN(x22, x11, 11) sqrN(x22, x11, 11, w)
mul(x22, x22, x11) mul(x22, x22, x11, w)
sqrN(x44, x22, 22) sqrN(x44, x22, 22, w)
mul(x44, x44, x22) mul(x44, x44, x22, w)
sqrN(x88, x44, 44) sqrN(x88, x44, 44, w)
mul(x88, x88, x44) mul(x88, x88, x44, w)
sqrN(x176, x88, 88) sqrN(x176, x88, 88, w)
mul(x176, x176, x88) mul(x176, x176, x88, w)
sqrN(x220, x176, 44) sqrN(x220, x176, 44, w)
mul(x220, x220, x44) mul(x220, x220, x44, w)
sqrN(x223, x220, 3) sqrN(x223, x220, 3, w)
mul(x223, x223, x3) mul(x223, x223, x3, w)
sqrN(out, x223, 23) sqrN(out, x223, 23, w)
mul(out, out, x22) mul(out, out, x22, w)
sqrN(out, out, 5) sqrN(out, out, 5, w)
mul(out, out, a) mul(out, out, a, w)
sqrN(out, out, 3) sqrN(out, out, 3, w)
mul(out, out, x2) mul(out, out, x2, w)
sqrN(out, out, 2) sqrN(out, out, 2, w)
mul(out, out, a) mul(out, out, a, w)
} }
fun sqrt( fun sqrt(
out: LongArray, out: LongArray,
a: LongArray, a: LongArray,
): Boolean { ): Boolean {
val w = wide.get()
val x2 = LongArray(4) val x2 = LongArray(4)
val x3 = LongArray(4) val x3 = LongArray(4)
val x6 = LongArray(4) val x6 = LongArray(4)
@@ -192,43 +223,53 @@ internal object FieldP {
val x220 = LongArray(4) val x220 = LongArray(4)
val x223 = LongArray(4) val x223 = LongArray(4)
sqr(x2, a) sqr(x2, a, w)
mul(x2, x2, a) mul(x2, x2, a, w)
sqr(x3, x2) sqr(x3, x2, w)
mul(x3, x3, a) mul(x3, x3, a, w)
sqrN(x6, x3, 3) sqrN(x6, x3, 3, w)
mul(x6, x6, x3) mul(x6, x6, x3, w)
sqrN(x9, x6, 3) sqrN(x9, x6, 3, w)
mul(x9, x9, x3) mul(x9, x9, x3, w)
sqrN(x11, x9, 2) sqrN(x11, x9, 2, w)
mul(x11, x11, x2) mul(x11, x11, x2, w)
sqrN(x22, x11, 11) sqrN(x22, x11, 11, w)
mul(x22, x22, x11) mul(x22, x22, x11, w)
sqrN(x44, x22, 22) sqrN(x44, x22, 22, w)
mul(x44, x44, x22) mul(x44, x44, x22, w)
sqrN(x88, x44, 44) sqrN(x88, x44, 44, w)
mul(x88, x88, x44) mul(x88, x88, x44, w)
sqrN(x176, x88, 88) sqrN(x176, x88, 88, w)
mul(x176, x176, x88) mul(x176, x176, x88, w)
sqrN(x220, x176, 44) sqrN(x220, x176, 44, w)
mul(x220, x220, x44) mul(x220, x220, x44, w)
sqrN(x223, x220, 3) sqrN(x223, x220, 3, w)
mul(x223, x223, x3) mul(x223, x223, x3, w)
sqrN(out, x223, 23) sqrN(out, x223, 23, w)
mul(out, out, x22) mul(out, out, x22, w)
sqrN(out, out, 6) sqrN(out, out, 6, w)
mul(out, out, x2) mul(out, out, x2, w)
sqrN(out, out, 2) sqrN(out, out, 2, w)
val check = LongArray(4) val check = LongArray(4)
mul(check, out, out) mul(check, out, out, w)
val ar = LongArray(4) val ar = LongArray(4)
U256.copyInto(ar, a) U256.copyInto(ar, a)
reduceSelf(ar) reduceSelf(ar)
return U256.cmp(check, ar) == 0 return U256.cmp(check, ar) == 0
} }
private fun sqrN(
out: LongArray,
a: LongArray,
n: Int,
w: LongArray,
) {
U256.copyInto(out, a)
repeat(n) { sqr(out, out, w) }
}
private fun sqrN( private fun sqrN(
out: LongArray, out: LongArray,
a: LongArray, a: LongArray,
@@ -242,18 +242,27 @@ internal object ECPoint {
/** /**
* Scratch space for point operations. Each thread gets its own set of temporary * Scratch space for point operations. Each thread gets its own set of temporary
* field elements to avoid allocation in the inner loops. The 12 temp buffers * field elements and a wide buffer to avoid allocation and ThreadLocal lookups
* (t[0]..t[11]) are shared across doublePoint and addPoints — this is safe because * in the inner loops. The 12 temp buffers (t[0]..t[11]) are shared across
* these functions only call each other in the equal-point degenerate case, which * doublePoint and addPoints — this is safe because these functions only call
* returns immediately after the recursive call without using the temps further. * each other in the equal-point degenerate case, which returns immediately
* after the recursive call without using the temps further.
*
* The wide buffer (LongArray(8)) is pre-fetched once per top-level operation
* and passed through to FieldP.mul/sqr, avoiding ~500+ ThreadLocal.get() calls
* per scalar multiplication (~20-30ns each on JVM).
*/ */
private class PointScratch { internal class PointScratch {
val t = Array(12) { LongArray(4) } val t = Array(12) { LongArray(4) }
val dblCopy = MutablePoint() // Copy buffer for in-place doubling (out === input) val dblCopy = MutablePoint() // Copy buffer for in-place doubling (out === input)
val w = LongArray(8) // Wide buffer for FieldP.mul/sqr — shared, avoids ThreadLocal
} }
private val scratch = ThreadLocal.withInitial { PointScratch() } private val scratch = ThreadLocal.withInitial { 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 (3M + 4S) ====================
/** /**
@@ -269,15 +278,22 @@ internal object ECPoint {
* *
* Safe for out === inp (in-place doubling) via internal copy buffer. * Safe for out === inp (in-place doubling) via internal copy buffer.
*/ */
/** doublePoint with ThreadLocal scratch (convenience for non-hot paths). */
fun doublePoint( fun doublePoint(
out: MutablePoint, out: MutablePoint,
inp: MutablePoint, inp: MutablePoint,
) = doublePoint(out, inp, scratch.get())
/** doublePoint with caller-provided scratch (hot path — no ThreadLocal lookup). */
fun doublePoint(
out: MutablePoint,
inp: MutablePoint,
s: PointScratch,
) { ) {
if (inp.isInfinity()) { if (inp.isInfinity()) {
out.setInfinity() out.setInfinity()
return return
} }
val s = scratch.get()
val p = val p =
if (out === inp) { if (out === inp) {
s.dblCopy.copyFrom(inp) s.dblCopy.copyFrom(inp)
@@ -286,23 +302,24 @@ internal object ECPoint {
inp inp
} }
val t = s.t val t = s.t
val w = s.w
FieldP.sqr(t[0], p.y) // S = Y² FieldP.sqr(t[0], p.y, w) // S = Y²
FieldP.sqr(t[1], p.x) // X² FieldP.sqr(t[1], p.x, w) // X²
FieldP.add(t[2], t[1], t[1]) // 2·X² FieldP.add(t[2], t[1], t[1]) // 2·X²
FieldP.add(t[2], t[2], t[1]) // 3·X² FieldP.add(t[2], t[2], t[1]) // 3·X²
FieldP.half(t[2], t[2]) // L = (3/2)·X² FieldP.half(t[2], t[2]) // L = (3/2)·X²
FieldP.mul(t[3], p.x, t[0]) // X·S FieldP.mul(t[3], p.x, t[0], w) // X·S
FieldP.neg(t[3], t[3]) // T = -X·S FieldP.neg(t[3], t[3]) // T = -X·S
FieldP.sqr(out.x, t[2]) // X₃ = L² FieldP.sqr(out.x, t[2], w) // X₃ = L²
FieldP.add(out.x, out.x, t[3]) // + T FieldP.add(out.x, out.x, t[3]) // + T
FieldP.add(out.x, out.x, t[3]) // + T FieldP.add(out.x, out.x, t[3]) // + T
FieldP.add(t[4], out.x, t[3]) // X₃ + T FieldP.add(t[4], out.x, t[3]) // X₃ + T
FieldP.mul(t[4], t[2], t[4]) // L·(X₃+T) FieldP.mul(t[4], t[2], t[4], w) // L·(X₃+T)
FieldP.sqr(t[5], t[0]) // S² FieldP.sqr(t[5], t[0], w) // S²
FieldP.add(t[4], t[4], t[5]) // L·(X₃+T) + S² FieldP.add(t[4], t[4], t[5]) // L·(X₃+T) + S²
FieldP.neg(out.y, t[4]) // Y₃ = negate FieldP.neg(out.y, t[4]) // Y₃ = negate
FieldP.mul(out.z, p.y, p.z) // Z₃ = Y·Z FieldP.mul(out.z, p.y, p.z, w) // Z₃ = Y·Z
} }
// ==================== Mixed Addition: Jacobian + Affine (8M + 3S) ==================== // ==================== Mixed Addition: Jacobian + Affine (8M + 3S) ====================
@@ -318,49 +335,58 @@ internal object ECPoint {
* *
* Handles degenerate cases: p is infinity, or p equals/negates q. * Handles degenerate cases: p is infinity, or p equals/negates q.
*/ */
/** addMixed with ThreadLocal scratch (convenience for non-hot paths). */
fun addMixed( fun addMixed(
out: MutablePoint, out: MutablePoint,
p: MutablePoint, p: MutablePoint,
qx: LongArray, qx: LongArray,
qy: LongArray, qy: LongArray,
) = addMixed(out, p, qx, qy, scratch.get())
/** addMixed with caller-provided scratch (hot path — no ThreadLocal lookup). */
fun addMixed(
out: MutablePoint,
p: MutablePoint,
qx: LongArray,
qy: LongArray,
s: PointScratch,
) { ) {
if (p.isInfinity()) { if (p.isInfinity()) {
out.setAffine(qx, qy) out.setAffine(qx, qy)
return return
} }
val s = scratch.get()
val t = s.t val t = s.t
val w = s.w
FieldP.sqr(t[0], p.z) // Z₁² FieldP.sqr(t[0], p.z, w) // Z₁²
FieldP.mul(t[1], t[0], p.z) // Z₁³ FieldP.mul(t[1], t[0], p.z, w) // Z₁³
FieldP.mul(t[2], qx, t[0]) // U₂ = qx·Z₁² (U₁ = X₁ since Z₂=1) FieldP.mul(t[2], qx, t[0], w) // U₂ = qx·Z₁² (U₁ = X₁ since Z₂=1)
FieldP.mul(t[3], qy, t[1]) // S₂ = qy·Z₁³ (S₁ = Y₁ since Z₂=1) FieldP.mul(t[3], qy, t[1], w) // S₂ = qy·Z₁³ (S₁ = Y₁ since Z₂=1)
FieldP.sub(t[4], t[2], p.x) // H = U₂ - U₁ FieldP.sub(t[4], t[2], p.x) // H = U₂ - U₁
if (U256.isZero(t[4])) { if (U256.isZero(t[4])) {
// Same x-coordinate: either same point (double) or inverse (infinity)
val tmp = LongArray(4) val tmp = LongArray(4)
FieldP.sub(tmp, t[3], p.y) FieldP.sub(tmp, t[3], p.y)
if (U256.isZero(tmp)) doublePoint(out, p) else out.setInfinity() if (U256.isZero(tmp)) doublePoint(out, p, s) else out.setInfinity()
return return
} }
FieldP.add(t[5], t[4], t[4]) // 2H FieldP.add(t[5], t[4], t[4]) // 2H
FieldP.sqr(t[5], t[5]) // I = (2H)² FieldP.sqr(t[5], t[5], w) // I = (2H)²
FieldP.mul(t[6], t[4], t[5]) // J = H·I FieldP.mul(t[6], t[4], t[5], w) // J = H·I
FieldP.sub(t[7], t[3], p.y) FieldP.sub(t[7], t[3], p.y)
FieldP.add(t[7], t[7], t[7]) // r = 2·(S₂ - S₁) FieldP.add(t[7], t[7], t[7]) // r = 2·(S₂ - S₁)
FieldP.mul(t[8], p.x, t[5]) // V = U₁·I FieldP.mul(t[8], p.x, t[5], w) // V = U₁·I
FieldP.sqr(out.x, t[7]) // X₃ = r² FieldP.sqr(out.x, t[7], w) // X₃ = r²
FieldP.sub(out.x, out.x, t[6]) // - J FieldP.sub(out.x, out.x, t[6]) // - J
FieldP.sub(out.x, out.x, t[8]) // - V FieldP.sub(out.x, out.x, t[8]) // - V
FieldP.sub(out.x, out.x, t[8]) // - V FieldP.sub(out.x, out.x, t[8]) // - V
FieldP.sub(t[9], t[8], out.x) // V - X₃ FieldP.sub(t[9], t[8], out.x) // V - X₃
FieldP.mul(out.y, t[7], t[9]) // Y₃ = r·(V-X₃) FieldP.mul(out.y, t[7], t[9], w) // Y₃ = r·(V-X₃)
FieldP.mul(t[9], p.y, t[6]) // - 2·S₁·J FieldP.mul(t[9], p.y, t[6], w) // - 2·S₁·J
FieldP.add(t[9], t[9], t[9]) FieldP.add(t[9], t[9], t[9])
FieldP.sub(out.y, out.y, t[9]) FieldP.sub(out.y, out.y, t[9])
FieldP.mul(out.z, p.z, t[4]) // Z₃ = 2·Z₁·H FieldP.mul(out.z, p.z, t[4], w) // Z₃ = 2·Z₁·H
FieldP.add(out.z, out.z, out.z) FieldP.add(out.z, out.z, out.z)
} }
@@ -378,6 +404,13 @@ internal object ECPoint {
out: MutablePoint, out: MutablePoint,
p: MutablePoint, p: MutablePoint,
q: MutablePoint, q: MutablePoint,
) = addPoints(out, p, q, scratch.get())
fun addPoints(
out: MutablePoint,
p: MutablePoint,
q: MutablePoint,
s: PointScratch,
) { ) {
if (p.isInfinity()) { if (p.isInfinity()) {
out.copyFrom(q) out.copyFrom(q)
@@ -387,45 +420,45 @@ internal object ECPoint {
out.copyFrom(p) out.copyFrom(p)
return return
} }
val s = scratch.get()
val t = s.t val t = s.t
val w = s.w
FieldP.sqr(t[0], p.z) // Z₁² FieldP.sqr(t[0], p.z, w) // Z₁²
FieldP.sqr(t[1], q.z) // Z₂² FieldP.sqr(t[1], q.z, w) // Z₂²
FieldP.mul(t[2], p.x, t[1]) // U₁ = X₁·Z₂² FieldP.mul(t[2], p.x, t[1], w) // U₁ = X₁·Z₂²
FieldP.mul(t[3], q.x, t[0]) // U₂ = X₂·Z₁² FieldP.mul(t[3], q.x, t[0], w) // U₂ = X₂·Z₁²
FieldP.mul(t[4], q.z, t[1]) // Z₂³ FieldP.mul(t[4], q.z, t[1], w) // Z₂³
FieldP.mul(t[4], p.y, t[4]) // S₁ = Y₁·Z₂³ FieldP.mul(t[4], p.y, t[4], w) // S₁ = Y₁·Z₂³
FieldP.mul(t[5], p.z, t[0]) // Z₁³ FieldP.mul(t[5], p.z, t[0], w) // Z₁³
FieldP.mul(t[5], q.y, t[5]) // S₂ = Y₂·Z₁³ FieldP.mul(t[5], q.y, t[5], w) // S₂ = Y₂·Z₁³
if (U256.cmp(t[2], t[3]) == 0) { if (U256.cmp(t[2], t[3]) == 0) {
if (U256.cmp(t[4], t[5]) == 0) doublePoint(out, p) else out.setInfinity() if (U256.cmp(t[4], t[5]) == 0) doublePoint(out, p, s) else out.setInfinity()
return return
} }
FieldP.sub(t[6], t[3], t[2]) // H = U₂ - U₁ FieldP.sub(t[6], t[3], t[2]) // H = U₂ - U₁
FieldP.add(t[7], t[6], t[6]) FieldP.add(t[7], t[6], t[6])
FieldP.sqr(t[7], t[7]) // I = (2H)² FieldP.sqr(t[7], t[7], w) // I = (2H)²
FieldP.mul(t[8], t[6], t[7]) // J = H·I FieldP.mul(t[8], t[6], t[7], w) // J = H·I
FieldP.sub(t[9], t[5], t[4]) FieldP.sub(t[9], t[5], t[4])
FieldP.add(t[9], t[9], t[9]) // r = 2·(S₂-S₁) FieldP.add(t[9], t[9], t[9]) // r = 2·(S₂-S₁)
FieldP.mul(t[10], t[2], t[7]) // V = U₁·I FieldP.mul(t[10], t[2], t[7], w) // V = U₁·I
FieldP.sqr(out.x, t[9]) FieldP.sqr(out.x, t[9], w)
FieldP.sub(out.x, out.x, t[8]) FieldP.sub(out.x, out.x, t[8])
FieldP.sub(out.x, out.x, t[10]) FieldP.sub(out.x, out.x, t[10])
FieldP.sub(out.x, out.x, t[10]) FieldP.sub(out.x, out.x, t[10])
FieldP.sub(t[11], t[10], out.x) FieldP.sub(t[11], t[10], out.x)
FieldP.mul(out.y, t[9], t[11]) FieldP.mul(out.y, t[9], t[11], w)
FieldP.mul(t[11], t[4], t[8]) FieldP.mul(t[11], t[4], t[8], w)
FieldP.add(t[11], t[11], t[11]) FieldP.add(t[11], t[11], t[11])
FieldP.sub(out.y, out.y, t[11]) FieldP.sub(out.y, out.y, t[11])
FieldP.add(out.z, p.z, q.z) FieldP.add(out.z, p.z, q.z)
FieldP.sqr(out.z, out.z) FieldP.sqr(out.z, out.z, w)
FieldP.sub(out.z, out.z, t[0]) FieldP.sub(out.z, out.z, t[0])
FieldP.sub(out.z, out.z, t[1]) FieldP.sub(out.z, out.z, t[1])
FieldP.mul(out.z, out.z, t[6]) FieldP.mul(out.z, out.z, t[6], w)
} }
// ==================== Scalar Multiplication ==================== // ==================== Scalar Multiplication ====================
@@ -449,26 +482,27 @@ internal object ECPoint {
return return
} }
val w = 5 val s = scratch.get()
val tableSize = 1 shl (w - 2) // 8 entries val wnd = 5
val tableSize = 1 shl (wnd - 2) // 8 entries
// Split scalar via GLV: scalar = k₁ + k₂·λ // Split scalar via GLV: scalar = k₁ + k₂·λ
val split = Glv.splitScalar(scalar) val split = Glv.splitScalar(scalar)
val wnaf1 = Glv.wnaf(split.k1, w, 129) val wnaf1 = Glv.wnaf(split.k1, wnd, 129)
val wnaf2 = Glv.wnaf(split.k2, w, 129) val wnaf2 = Glv.wnaf(split.k2, wnd, 129)
// P odd-multiples [1P, 3P, 5P, ..., 15P] via 2P stepping // P odd-multiples [1P, 3P, 5P, ..., 15P] via 2P stepping
val p2 = MutablePoint() val p2 = MutablePoint()
doublePoint(p2, p) doublePoint(p2, p, s)
val pOdd = Array(tableSize) { MutablePoint() } val pOdd = Array(tableSize) { MutablePoint() }
pOdd[0].copyFrom(p) pOdd[0].copyFrom(p)
for (i in 1 until tableSize) addPoints(pOdd[i], pOdd[i - 1], p2) for (i in 1 until tableSize) addPoints(pOdd[i], pOdd[i - 1], p2, s)
// λ(P) odd-multiples: (β·X, Y, Z) in Jacobian // λ(P) odd-multiples: (β·X, Y, Z) in Jacobian
val pLamOdd = val pLamOdd =
Array(tableSize) { i -> Array(tableSize) { i ->
val lp = MutablePoint() val lp = MutablePoint()
FieldP.mul(lp.x, pOdd[i].x, Glv.BETA) FieldP.mul(lp.x, pOdd[i].x, Glv.BETA, s.w)
pOdd[i].y.copyInto(lp.y) pOdd[i].y.copyInto(lp.y)
pOdd[i].z.copyInto(lp.z) pOdd[i].z.copyInto(lp.z)
lp lp
@@ -487,9 +521,9 @@ internal object ECPoint {
val negJac = MutablePoint() val negJac = MutablePoint()
for (i in bits - 1 downTo 0) { for (i in bits - 1 downTo 0) {
doublePoint(out, out) doublePoint(out, out, s)
addWnafJacobian(out, tmp, negJac, wnaf1, i, pOdd, split.negK1) addWnafJacobian(out, tmp, negJac, wnaf1, i, pOdd, split.negK1, s)
addWnafJacobian(out, tmp, negJac, wnaf2, i, pLamOdd, split.negK2) addWnafJacobian(out, tmp, negJac, wnaf2, i, pLamOdd, split.negK2, s)
} }
} }
@@ -512,13 +546,14 @@ internal object ECPoint {
return return
} }
val s = scratch.get()
val table = combTable val table = combTable
out.setInfinity() out.setInfinity()
val tmp = MutablePoint() val tmp = MutablePoint()
for (combOff in COMB_SPACING - 1 downTo 0) { for (combOff in COMB_SPACING - 1 downTo 0) {
if (combOff < COMB_SPACING - 1) { if (combOff < COMB_SPACING - 1) {
doublePoint(out, out) doublePoint(out, out, s)
} }
for (block in 0 until COMB_BLOCKS) { for (block in 0 until COMB_BLOCKS) {
var mask = 0 var mask = 0
@@ -530,7 +565,7 @@ internal object ECPoint {
} }
if (mask != 0) { if (mask != 0) {
val entry = table[block * COMB_POINTS + mask] val entry = table[block * COMB_POINTS + mask]
addMixed(tmp, out, entry.x, entry.y) addMixed(tmp, out, entry.x, entry.y, s)
out.copyFrom(tmp) out.copyFrom(tmp)
} }
} }
@@ -555,6 +590,7 @@ internal object ECPoint {
p: MutablePoint, p: MutablePoint,
e: LongArray, e: LongArray,
) { ) {
val sc = scratch.get()
val wP = 5 // Window for P-side (table built per-call, keep small) val wP = 5 // Window for P-side (table built per-call, keep small)
val pTableSize = 1 shl (wP - 2) // 8 entries for P val pTableSize = 1 shl (wP - 2) // 8 entries for P
@@ -574,15 +610,15 @@ internal object ECPoint {
// P odd-multiples [1P, 3P, 5P, ..., 15P] via 2P stepping (1 double + 7 adds) // P odd-multiples [1P, 3P, 5P, ..., 15P] via 2P stepping (1 double + 7 adds)
val p2 = MutablePoint() val p2 = MutablePoint()
doublePoint(p2, p) doublePoint(p2, p, sc)
val pOdd = Array(pTableSize) { MutablePoint() } val pOdd = Array(pTableSize) { MutablePoint() }
pOdd[0].copyFrom(p) pOdd[0].copyFrom(p)
for (i in 1 until pTableSize) addPoints(pOdd[i], pOdd[i - 1], p2) for (i in 1 until pTableSize) addPoints(pOdd[i], pOdd[i - 1], p2, sc)
// λ(P) table: (β·X, Y, Z) in Jacobian — endomorphism preserves projective coords // λ(P) table: (β·X, Y, Z) in Jacobian — endomorphism preserves projective coords
val pLamOdd = val pLamOdd =
Array(pTableSize) { i -> Array(pTableSize) { i ->
val lp = MutablePoint() val lp = MutablePoint()
FieldP.mul(lp.x, pOdd[i].x, Glv.BETA) FieldP.mul(lp.x, pOdd[i].x, Glv.BETA, sc.w)
pOdd[i].y.copyInto(lp.y) pOdd[i].y.copyInto(lp.y)
pOdd[i].z.copyInto(lp.z) pOdd[i].z.copyInto(lp.z)
lp lp
@@ -600,16 +636,16 @@ internal object ECPoint {
out.setInfinity() out.setInfinity()
val tmp = MutablePoint() val tmp = MutablePoint()
val negY = LongArray(4) val negY = LongArray(4)
val negJac = MutablePoint() // Reused scratch for Jacobian negation val negJac = MutablePoint()
for (i in bits - 1 downTo 0) { for (i in bits - 1 downTo 0) {
doublePoint(out, out) doublePoint(out, out, sc)
// Streams 1-2: G-side (affine tables, mixed addition) // Streams 1-2: G-side (affine tables, mixed addition)
addWnafMixed(out, tmp, negY, wnafS1, i, gOdd, sSplit.negK1) addWnafMixed(out, tmp, negY, wnafS1, i, gOdd, sSplit.negK1, sc)
addWnafMixed(out, tmp, negY, wnafS2, i, gLam, sSplit.negK2) addWnafMixed(out, tmp, negY, wnafS2, i, gLam, sSplit.negK2, sc)
// Streams 3-4: P-side (Jacobian tables, full addition) // Streams 3-4: P-side (Jacobian tables, full addition)
addWnafJacobian(out, tmp, negJac, wnafE1, i, pOdd, eSplit.negK1) addWnafJacobian(out, tmp, negJac, wnafE1, i, pOdd, eSplit.negK1, sc)
addWnafJacobian(out, tmp, negJac, wnafE2, i, pLamOdd, eSplit.negK2) addWnafJacobian(out, tmp, negJac, wnafE2, i, pLamOdd, eSplit.negK2, sc)
} }
} }
@@ -626,6 +662,7 @@ internal object ECPoint {
bitIndex: Int, bitIndex: Int,
table: Array<AffinePoint>, table: Array<AffinePoint>,
glvNeg: Boolean, glvNeg: Boolean,
s: PointScratch,
) { ) {
if (bitIndex >= wnafDigits.size) return if (bitIndex >= wnafDigits.size) return
val d = wnafDigits[bitIndex] val d = wnafDigits[bitIndex]
@@ -633,10 +670,10 @@ internal object ECPoint {
val idx = (if (d > 0) d else -d) / 2 val idx = (if (d > 0) d else -d) / 2
val effectiveNeg = (d < 0) xor glvNeg val effectiveNeg = (d < 0) xor glvNeg
if (!effectiveNeg) { if (!effectiveNeg) {
addMixed(tmp, out, table[idx].x, table[idx].y) addMixed(tmp, out, table[idx].x, table[idx].y, s)
} else { } else {
FieldP.neg(negY, table[idx].y) FieldP.neg(negY, table[idx].y)
addMixed(tmp, out, table[idx].x, negY) addMixed(tmp, out, table[idx].x, negY, s)
} }
out.copyFrom(tmp) out.copyFrom(tmp)
} }
@@ -650,6 +687,7 @@ internal object ECPoint {
bitIndex: Int, bitIndex: Int,
table: Array<MutablePoint>, table: Array<MutablePoint>,
glvNeg: Boolean, glvNeg: Boolean,
s: PointScratch,
) { ) {
if (bitIndex >= wnafDigits.size) return if (bitIndex >= wnafDigits.size) return
val d = wnafDigits[bitIndex] val d = wnafDigits[bitIndex]
@@ -657,13 +695,12 @@ internal object ECPoint {
val idx = (if (d > 0) d else -d) / 2 val idx = (if (d > 0) d else -d) / 2
val effectiveNeg = (d < 0) xor glvNeg val effectiveNeg = (d < 0) xor glvNeg
if (!effectiveNeg) { if (!effectiveNeg) {
addPoints(tmp, out, table[idx]) addPoints(tmp, out, table[idx], s)
} else { } else {
// Negate the Jacobian point: (X, -Y, Z) using pre-allocated scratch
table[idx].x.copyInto(negScratch.x) table[idx].x.copyInto(negScratch.x)
FieldP.neg(negScratch.y, table[idx].y) FieldP.neg(negScratch.y, table[idx].y)
table[idx].z.copyInto(negScratch.z) table[idx].z.copyInto(negScratch.z)
addPoints(tmp, out, negScratch) addPoints(tmp, out, negScratch, s)
} }
out.copyFrom(tmp) out.copyFrom(tmp)
} }