diff --git a/marmotBench/README.md b/marmotBench/README.md index 8e6b59294c..0d151b2d57 100644 --- a/marmotBench/README.md +++ b/marmotBench/README.md @@ -60,3 +60,51 @@ landing inside a measured sample and corrupting the percentile it falls in. of *this* workload on *this* machine. It is not a language benchmark. - `create_group/32` builds 32 KeyPackages in setup. That cost is excluded, but it makes each iteration expensive to prepare — hence the low iteration count. + +## Result: eliminating the field-arithmetic allocation + +The first run of this module put **93% of all sampled allocation** (JFR +`jdk.ObjectAllocationSample`) in `Curve25519Field.mul/add/sub`. The pure-Kotlin +Curve25519 returned a fresh `LongArray(16)` from every field operation, and a +Montgomery ladder performs ~18 of them per bit for 255 bits — so a single +X25519 scalar multiplication allocated over a megabyte of garbage. + +Each operation now has an in-place `*Into` twin, and both hot paths (the X25519 +ladder and Ed25519's extended-coordinate point addition) allocate their working +set once and then run allocation-free. See `Curve25519Field`. + +Allocation per operation, before and after. This column reproduces to four +significant figures across runs, so the ratios are real: + +| operation | before | after | reduction | +|---------------------|--------------|-------------|-----------| +| `create_group/0` | 6 958.7 KB | 88.3 KB | 79x | +| `create_group/1` | 27 074.1 KB | 570.5 KB | 47x | +| `create_group/8` | 94 548.6 KB | 3 501.7 KB | 27x | +| `create_group/32` | 333 101.2 KB | 33 216.6 KB | 10x | +| `join_welcome` | 10 331.6 KB | 252.2 KB | 41x | +| `send_app_message` | 2 755.1 KB | 71.6 KB | 38x | +| `ingest_app_message`| 5 640.3 KB | 70.9 KB | 80x | + +Latency improved too, though it is the noisier measurement — two post-rewrite +runs are given so the spread is visible rather than averaged away: + +| operation | p50 before | p50 after (run 1 / run 2) | +|---------------------|------------|---------------------------| +| `create_group/0` | 6 323.7us | 4 184.2 / 3 971.5us | +| `create_group/1` | 16 928.2us | 13 197.2 / 13 574.2us | +| `create_group/8` | 51 470.7us | 43 244.0 / 43 424.7us | +| `create_group/32` | 190 080.0us | 202 261.1 / 172 908.9us | +| `join_welcome` | 6 219.3us | 4 919.5 / 4 991.3us | +| `send_app_message` | 1 720.9us | 1 314.3 / 1 376.0us | +| `ingest_app_message`| 3 114.1us | 2 487.5 / 2 545.6us | + +`create_group/32` is the row to distrust: it has the fewest iterations, and its +two runs disagree by 17% at p50 and by nearly 2x at p99 (400.3ms then 210.4ms). +Read it as "no worse"; the other rows are consistent enough to read as gains. + +Against MDK this closes most of the `create_group` gap — `create_group/1` goes +from 4.7x slower to about 3.7x — without changing a single protocol behaviour: +the RFC 7748 / RFC 8032 vector suites, the HPKE tests and the full 4833-test +quartz suite all pass unchanged, which is the point of keeping the allocating +functions around to differentially test against. diff --git a/quartz/src/appleMain/kotlin/com/vitorpamplona/quartz/marmot/mls/crypto/Ed25519.apple.kt b/quartz/src/appleMain/kotlin/com/vitorpamplona/quartz/marmot/mls/crypto/Ed25519.apple.kt index 629e54a01d..6f6d950d46 100644 --- a/quartz/src/appleMain/kotlin/com/vitorpamplona/quartz/marmot/mls/crypto/Ed25519.apple.kt +++ b/quartz/src/appleMain/kotlin/com/vitorpamplona/quartz/marmot/mls/crypto/Ed25519.apple.kt @@ -179,34 +179,69 @@ actual object Ed25519 { p[2].copyOf(), p[3].copyOf(), ) - addPointInPlace(result, q) + addPointInPlace(result, q, PointAddScratch()) return result } + /** + * Scratch for [addPointInPlace]. + * + * Extended-coordinate addition needs ten temporary field elements, and a + * scalar multiplication calls it 512 times — twice per scalar bit. Making + * each call allocate its own was the single largest source of garbage in + * the MLS stack: an allocation profile put 93% of all sampled allocation + * in `Curve25519Field.mul/add/sub`, most of it reached from here. + * + * One instance is created per scalar multiplication and reused by every + * step, so the loop allocates nothing. + */ + private class PointAddScratch { + val a = LongArray(16) + val b = LongArray(16) + val c = LongArray(16) + val d = LongArray(16) + val e = LongArray(16) + val f = LongArray(16) + val g = LongArray(16) + val h = LongArray(16) + val t1 = LongArray(16) + val t2 = LongArray(16) + + /** The 31-limb accumulator every [Curve25519Field.mulInto] here shares. */ + val mulT = LongArray(31) + } + /** In-place point addition: p += q. */ private fun addPointInPlace( p: Array, q: Array, + s: PointAddScratch, ) { - val a = Curve25519Field.sub(p[1], p[0]) - val t = Curve25519Field.sub(q[1], q[0]) - val aMul = Curve25519Field.mul(a, t) - val b = Curve25519Field.add(p[0], p[1]) - val t2 = Curve25519Field.add(q[0], q[1]) - val bMul = Curve25519Field.mul(b, t2) - val c = Curve25519Field.mul(p[3], q[3]) - val cMul = Curve25519Field.mul(c, Curve25519Field.D2) - val d = Curve25519Field.mul(p[2], q[2]) - val dAdd = Curve25519Field.add(d, d) - val e = Curve25519Field.sub(bMul, aMul) - val f = Curve25519Field.sub(dAdd, cMul) - val g = Curve25519Field.add(dAdd, cMul) - val h = Curve25519Field.add(bMul, aMul) + // Every read of p and q happens in this first block. `scalarMult` + // doubles by passing the same point as both arguments, so nothing may + // be written back until these are done. + Curve25519Field.subInto(s.a, p[1], p[0]) + Curve25519Field.subInto(s.t1, q[1], q[0]) + Curve25519Field.mulInto(s.a, s.a, s.t1, s.mulT) + Curve25519Field.addInto(s.b, p[0], p[1]) + Curve25519Field.addInto(s.t2, q[0], q[1]) + Curve25519Field.mulInto(s.b, s.b, s.t2, s.mulT) + Curve25519Field.mulInto(s.c, p[3], q[3], s.mulT) + Curve25519Field.mulInto(s.c, s.c, Curve25519Field.D2, s.mulT) + Curve25519Field.mulInto(s.d, p[2], q[2], s.mulT) + Curve25519Field.addInto(s.d, s.d, s.d) - Curve25519Field.mul(e, f).copyInto(p[0]) - Curve25519Field.mul(h, g).copyInto(p[1]) - Curve25519Field.mul(g, f).copyInto(p[2]) - Curve25519Field.mul(e, h).copyInto(p[3]) + Curve25519Field.subInto(s.e, s.b, s.a) + Curve25519Field.subInto(s.f, s.d, s.c) + Curve25519Field.addInto(s.g, s.d, s.c) + Curve25519Field.addInto(s.h, s.b, s.a) + + // Safe to write p now: e, f, g and h are scratch, so no later product + // reads anything we are about to overwrite. + Curve25519Field.mulInto(p[0], s.e, s.f, s.mulT) + Curve25519Field.mulInto(p[1], s.h, s.g, s.mulT) + Curve25519Field.mulInto(p[2], s.g, s.f, s.mulT) + Curve25519Field.mulInto(p[3], s.e, s.h, s.mulT) } /** Point doubling (self-addition). */ @@ -235,11 +270,13 @@ actual object Ed25519 { p[2].copyOf(), p[3].copyOf(), ) + // One workspace for all 512 additions below. + val scratch = PointAddScratch() for (i in 255 downTo 0) { val b = ((s[i shr 3].toInt() shr (i and 7)) and 1).toLong() cswap(result, q, b) - addPointInPlace(q, result) - addPointInPlace(result, result) + addPointInPlace(q, result, scratch) + addPointInPlace(result, result, scratch) cswap(result, q, b) } return result diff --git a/quartz/src/appleMain/kotlin/com/vitorpamplona/quartz/marmot/mls/crypto/X25519.apple.kt b/quartz/src/appleMain/kotlin/com/vitorpamplona/quartz/marmot/mls/crypto/X25519.apple.kt index e08b01ccd9..1625e6b4f4 100644 --- a/quartz/src/appleMain/kotlin/com/vitorpamplona/quartz/marmot/mls/crypto/X25519.apple.kt +++ b/quartz/src/appleMain/kotlin/com/vitorpamplona/quartz/marmot/mls/crypto/X25519.apple.kt @@ -81,45 +81,62 @@ actual object X25519 { val c = Curve25519Field.GF0.copyOf() val d = Curve25519Field.GF1.copyOf() + // The ladder's entire working set, allocated ONCE. Every field + // operation in the loop writes into one of these, so 255 iterations + // allocate nothing at all — where the allocating form produced a fresh + // element per operation, about 1.3 MB of garbage per call. + val e = LongArray(16) + val f = LongArray(16) + val g = LongArray(16) + val h = LongArray(16) + val dd = LongArray(16) + val ff = LongArray(16) + val da = LongArray(16) + val cb = LongArray(16) + val cc = LongArray(16) + val tmp = LongArray(16) + val t = LongArray(31) + for (i in 254 downTo 0) { val r = ((z[i shr 3].toLong() shr (i and 7)) and 1) Curve25519Field.sel25519(a, b, r) Curve25519Field.sel25519(c, d, r) - val e = Curve25519Field.add(a, c) - val aMc = Curve25519Field.sub(a, c) - val f = Curve25519Field.add(b, d) - val bMd = Curve25519Field.sub(b, d) + // a, b, c and d are read only by these four lines; from here on + // they are dead and can be overwritten with the new values. + Curve25519Field.addInto(e, a, c) + Curve25519Field.subInto(g, a, c) + Curve25519Field.addInto(f, b, d) + Curve25519Field.subInto(h, b, d) - val dd = Curve25519Field.sqr(e) - val ff = Curve25519Field.sqr(aMc) - val da = Curve25519Field.mul(bMd, e) - val cb = Curve25519Field.mul(f, aMc) + Curve25519Field.sqrInto(dd, e, t) + Curve25519Field.sqrInto(ff, g, t) + Curve25519Field.mulInto(da, h, e, t) + Curve25519Field.mulInto(cb, f, g, t) - val ePrime = Curve25519Field.add(da, cb) - val aPrime = Curve25519Field.sub(da, cb) + // e := da + cb and g := da - cb. Reusing e and g is safe: both + // held inputs to the four products above, which are now computed. + Curve25519Field.addInto(e, da, cb) + Curve25519Field.subInto(g, da, cb) - val bNew = Curve25519Field.sqr(ePrime) - val aSqr = Curve25519Field.sqr(aPrime) - val dNew = Curve25519Field.mul(aSqr, x) + Curve25519Field.sqrInto(b, e, t) + Curve25519Field.sqrInto(g, g, t) + Curve25519Field.mulInto(d, g, x, t) - val aNew = Curve25519Field.mul(dd, ff) - val cc = Curve25519Field.sub(dd, ff) - val tmp = Curve25519Field.mul(cc, Curve25519Field.A24) - val ddPlusTmp = Curve25519Field.add(dd, tmp) - val cNew = Curve25519Field.mul(cc, ddPlusTmp) - - aNew.copyInto(a) - bNew.copyInto(b) - cNew.copyInto(c) - dNew.copyInto(d) + Curve25519Field.mulInto(a, dd, ff, t) + Curve25519Field.subInto(cc, dd, ff) + Curve25519Field.mulInto(tmp, cc, Curve25519Field.A24, t) + Curve25519Field.addInto(tmp, dd, tmp) + Curve25519Field.mulInto(c, cc, tmp, t) Curve25519Field.sel25519(a, b, r) Curve25519Field.sel25519(c, d, r) } - val invC = Curve25519Field.inv25519(c) - val result = Curve25519Field.mul(a, invC) - return Curve25519Field.pack25519(result) + // c := 1/c, then a := a/c. `tmp` is free again and serves as the + // inversion's scratch element. + Curve25519Field.inv25519Into(c, c, tmp, t) + Curve25519Field.mulInto(a, a, c, t) + return Curve25519Field.pack25519(a) } } diff --git a/quartz/src/commonMain/kotlin/com/vitorpamplona/quartz/marmot/mls/crypto/Curve25519Field.kt b/quartz/src/commonMain/kotlin/com/vitorpamplona/quartz/marmot/mls/crypto/Curve25519Field.kt index b36ad70326..45e05975f0 100644 --- a/quartz/src/commonMain/kotlin/com/vitorpamplona/quartz/marmot/mls/crypto/Curve25519Field.kt +++ b/quartz/src/commonMain/kotlin/com/vitorpamplona/quartz/marmot/mls/crypto/Curve25519Field.kt @@ -195,58 +195,139 @@ internal object Curve25519Field { return o } + // Allocating vs in-place. + // + // Each `add`/`sub`/`mul`/`sqr` below returns a NEW field element, which + // reads well and is what the TweetNaCl reference does. Inside a scalar + // multiplication it is also ~1.3 MB of garbage per call: a Montgomery + // ladder runs 255 iterations of ten muls and eight add/subs, and every one + // of them allocated. An allocation profile of the Marmot benchmarks put + // 93% of ALL sampled allocation in these three functions. + // + // So each one has an `*Into` twin that writes into a caller-owned output. + // The hot paths (X25519 and Ed25519 scalar multiplication) allocate + // their working set once and then run allocation-free. + // + // Both forms stay: the allocating ones are used off the hot path, where + // the clarity is worth more than the bytes, and keeping them means the + // in-place versions can be differentially tested against them. + // + // Every `*Into` is safe when the output aliases an input — the ladder + // relies on that. + /** Field addition: o = a + b. */ fun add( a: LongArray, b: LongArray, ): LongArray { val o = LongArray(16) - for (i in 0 until 16) o[i] = a[i] + b[i] + addInto(o, a, b) return o } + /** Field addition into [o]. Safe when [o] aliases [a] or [b]. */ + fun addInto( + o: LongArray, + a: LongArray, + b: LongArray, + ) { + for (i in 0 until 16) o[i] = a[i] + b[i] + } + /** Field subtraction: o = a - b. */ fun sub( a: LongArray, b: LongArray, ): LongArray { val o = LongArray(16) - for (i in 0 until 16) o[i] = a[i] - b[i] + subInto(o, a, b) return o } + /** Field subtraction into [o]. Safe when [o] aliases [a] or [b]. */ + fun subInto( + o: LongArray, + a: LongArray, + b: LongArray, + ) { + for (i in 0 until 16) o[i] = a[i] - b[i] + } + /** Field multiplication: o = a * b (mod p). */ fun mul( a: LongArray, b: LongArray, ): LongArray { - val t = LongArray(31) + val o = LongArray(16) + mulInto(o, a, b, LongArray(31)) + return o + } + + /** + * Field multiplication into [o], using [t] as the 31-limb accumulator. + * + * [t] is caller-owned so a loop can reuse one across thousands of calls; + * it is zeroed here, so callers never have to. Safe when [o] aliases [a] + * or [b]: the product is fully accumulated in [t] before [o] is touched. + */ + fun mulInto( + o: LongArray, + a: LongArray, + b: LongArray, + t: LongArray, + ) { + t.fill(0L) for (i in 0 until 16) { + val ai = a[i] for (j in 0 until 16) { - t[i + j] += a[i] * b[j] + t[i + j] += ai * b[j] } } for (i in 0 until 15) { t[i] += 38 * t[i + 16] } - val o = LongArray(16) for (i in 0 until 16) o[i] = t[i] car25519(o) car25519(o) - return o } /** Field squaring: o = a^2 (mod p). */ fun sqr(a: LongArray): LongArray = mul(a, a) + /** Field squaring into [o]. See [mulInto] for the [t] contract. */ + fun sqrInto( + o: LongArray, + a: LongArray, + t: LongArray, + ) = mulInto(o, a, a, t) + /** Field inversion: o = a^(-1) (mod p) using Fermat's little theorem. */ fun inv25519(a: LongArray): LongArray { - var c = a.copyOf() + val o = LongArray(16) + inv25519Into(o, a, LongArray(16), LongArray(31)) + return o + } + + /** + * Field inversion into [o], allocation-free. + * + * 254 squarings and ~250 multiplications, which is why this one matters: + * on the allocating path it was the single largest contributor after the + * ladder itself. [c] is a scratch field element and [t] the [mulInto] + * accumulator; [o] may alias [a]. + */ + fun inv25519Into( + o: LongArray, + a: LongArray, + c: LongArray, + t: LongArray, + ) { + a.copyInto(c) for (i in 253 downTo 0) { - c = sqr(c) - if (i != 2 && i != 4) c = mul(c, a) + sqrInto(c, c, t) + if (i != 2 && i != 4) mulInto(c, c, a, t) } - return c + c.copyInto(o) } /** Parity of a field element (lowest bit after reduction). */ @@ -257,10 +338,11 @@ internal object Curve25519Field { /** Raise a field element to the power (2^252 - 3), used in sqrt. */ fun pow2523(a: LongArray): LongArray { - var c = a.copyOf() + val c = a.copyOf() + val t = LongArray(31) for (i in 250 downTo 0) { - c = sqr(c) - if (i != 1) c = mul(c, a) + sqrInto(c, c, t) + if (i != 1) mulInto(c, c, a, t) } return c } diff --git a/quartz/src/jvmAndroid/kotlin/com/vitorpamplona/quartz/marmot/mls/crypto/Ed25519.jvmAndroid.kt b/quartz/src/jvmAndroid/kotlin/com/vitorpamplona/quartz/marmot/mls/crypto/Ed25519.jvmAndroid.kt index 752c490a56..028eddef0a 100644 --- a/quartz/src/jvmAndroid/kotlin/com/vitorpamplona/quartz/marmot/mls/crypto/Ed25519.jvmAndroid.kt +++ b/quartz/src/jvmAndroid/kotlin/com/vitorpamplona/quartz/marmot/mls/crypto/Ed25519.jvmAndroid.kt @@ -155,33 +155,68 @@ actual object Ed25519 { p[2].copyOf(), p[3].copyOf(), ) - addPointInPlace(result, q) + addPointInPlace(result, q, PointAddScratch()) return result } + /** + * Scratch for [addPointInPlace]. + * + * Extended-coordinate addition needs ten temporary field elements, and a + * scalar multiplication calls it 512 times — twice per scalar bit. Making + * each call allocate its own was the single largest source of garbage in + * the MLS stack: an allocation profile put 93% of all sampled allocation + * in `Curve25519Field.mul/add/sub`, most of it reached from here. + * + * One instance is created per scalar multiplication and reused by every + * step, so the loop allocates nothing. + */ + private class PointAddScratch { + val a = LongArray(16) + val b = LongArray(16) + val c = LongArray(16) + val d = LongArray(16) + val e = LongArray(16) + val f = LongArray(16) + val g = LongArray(16) + val h = LongArray(16) + val t1 = LongArray(16) + val t2 = LongArray(16) + + /** The 31-limb accumulator every [Curve25519Field.mulInto] here shares. */ + val mulT = LongArray(31) + } + private fun addPointInPlace( p: Array, q: Array, + s: PointAddScratch, ) { - val a = Curve25519Field.sub(p[1], p[0]) - val t = Curve25519Field.sub(q[1], q[0]) - val aMul = Curve25519Field.mul(a, t) - val b = Curve25519Field.add(p[0], p[1]) - val t2 = Curve25519Field.add(q[0], q[1]) - val bMul = Curve25519Field.mul(b, t2) - val c = Curve25519Field.mul(p[3], q[3]) - val cMul = Curve25519Field.mul(c, Curve25519Field.D2) - val d = Curve25519Field.mul(p[2], q[2]) - val dAdd = Curve25519Field.add(d, d) - val e = Curve25519Field.sub(bMul, aMul) - val f = Curve25519Field.sub(dAdd, cMul) - val g = Curve25519Field.add(dAdd, cMul) - val h = Curve25519Field.add(bMul, aMul) + // Every read of p and q happens in this first block. `scalarMult` + // doubles by passing the same point as both arguments, so nothing may + // be written back until these are done. + Curve25519Field.subInto(s.a, p[1], p[0]) + Curve25519Field.subInto(s.t1, q[1], q[0]) + Curve25519Field.mulInto(s.a, s.a, s.t1, s.mulT) + Curve25519Field.addInto(s.b, p[0], p[1]) + Curve25519Field.addInto(s.t2, q[0], q[1]) + Curve25519Field.mulInto(s.b, s.b, s.t2, s.mulT) + Curve25519Field.mulInto(s.c, p[3], q[3], s.mulT) + Curve25519Field.mulInto(s.c, s.c, Curve25519Field.D2, s.mulT) + Curve25519Field.mulInto(s.d, p[2], q[2], s.mulT) + Curve25519Field.addInto(s.d, s.d, s.d) - Curve25519Field.mul(e, f).copyInto(p[0]) - Curve25519Field.mul(h, g).copyInto(p[1]) - Curve25519Field.mul(g, f).copyInto(p[2]) - Curve25519Field.mul(e, h).copyInto(p[3]) + Curve25519Field.subInto(s.e, s.b, s.a) + Curve25519Field.subInto(s.f, s.d, s.c) + Curve25519Field.addInto(s.g, s.d, s.c) + Curve25519Field.addInto(s.h, s.b, s.a) + + // Safe to write p now: e, f, g and h are scratch, so no later product + // reads anything we are about to overwrite. + Curve25519Field.mulInto(p[0], s.e, s.f, s.mulT) + Curve25519Field.mulInto(p[1], s.h, s.g, s.mulT) + Curve25519Field.mulInto(p[2], s.g, s.f, s.mulT) + Curve25519Field.mulInto(p[3], s.e, s.h, s.mulT) } private fun negatePoint(p: Array): Array { @@ -205,11 +240,13 @@ actual object Ed25519 { p[2].copyOf(), p[3].copyOf(), ) + // One workspace for all 512 additions below. + val scratch = PointAddScratch() for (i in 255 downTo 0) { val b = ((s[i shr 3].toInt() shr (i and 7)) and 1).toLong() cswap(result, q, b) - addPointInPlace(q, result) - addPointInPlace(result, result) + addPointInPlace(q, result, scratch) + addPointInPlace(result, result, scratch) cswap(result, q, b) } return result diff --git a/quartz/src/jvmAndroid/kotlin/com/vitorpamplona/quartz/marmot/mls/crypto/X25519.jvmAndroid.kt b/quartz/src/jvmAndroid/kotlin/com/vitorpamplona/quartz/marmot/mls/crypto/X25519.jvmAndroid.kt index ab4b69b788..3caeab777f 100644 --- a/quartz/src/jvmAndroid/kotlin/com/vitorpamplona/quartz/marmot/mls/crypto/X25519.jvmAndroid.kt +++ b/quartz/src/jvmAndroid/kotlin/com/vitorpamplona/quartz/marmot/mls/crypto/X25519.jvmAndroid.kt @@ -83,45 +83,62 @@ actual object X25519 { val c = Curve25519Field.GF0.copyOf() val d = Curve25519Field.GF1.copyOf() + // The ladder's entire working set, allocated ONCE. Every field + // operation in the loop writes into one of these, so 255 iterations + // allocate nothing at all — where the allocating form produced a fresh + // element per operation, about 1.3 MB of garbage per call. + val e = LongArray(16) + val f = LongArray(16) + val g = LongArray(16) + val h = LongArray(16) + val dd = LongArray(16) + val ff = LongArray(16) + val da = LongArray(16) + val cb = LongArray(16) + val cc = LongArray(16) + val tmp = LongArray(16) + val t = LongArray(31) + for (i in 254 downTo 0) { val r = ((z[i shr 3].toLong() shr (i and 7)) and 1) Curve25519Field.sel25519(a, b, r) Curve25519Field.sel25519(c, d, r) - val e = Curve25519Field.add(a, c) - val aMc = Curve25519Field.sub(a, c) - val f = Curve25519Field.add(b, d) - val bMd = Curve25519Field.sub(b, d) + // a, b, c and d are read only by these four lines; from here on + // they are dead and can be overwritten with the new values. + Curve25519Field.addInto(e, a, c) + Curve25519Field.subInto(g, a, c) + Curve25519Field.addInto(f, b, d) + Curve25519Field.subInto(h, b, d) - val dd = Curve25519Field.sqr(e) - val ff = Curve25519Field.sqr(aMc) - val da = Curve25519Field.mul(bMd, e) - val cb = Curve25519Field.mul(f, aMc) + Curve25519Field.sqrInto(dd, e, t) + Curve25519Field.sqrInto(ff, g, t) + Curve25519Field.mulInto(da, h, e, t) + Curve25519Field.mulInto(cb, f, g, t) - val ePrime = Curve25519Field.add(da, cb) - val aPrime = Curve25519Field.sub(da, cb) + // e := da + cb and g := da - cb. Reusing e and g is safe: both + // held inputs to the four products above, which are now computed. + Curve25519Field.addInto(e, da, cb) + Curve25519Field.subInto(g, da, cb) - val bNew = Curve25519Field.sqr(ePrime) - val aSqr = Curve25519Field.sqr(aPrime) - val dNew = Curve25519Field.mul(aSqr, x) + Curve25519Field.sqrInto(b, e, t) + Curve25519Field.sqrInto(g, g, t) + Curve25519Field.mulInto(d, g, x, t) - val aNew = Curve25519Field.mul(dd, ff) - val cc = Curve25519Field.sub(dd, ff) - val tmp = Curve25519Field.mul(cc, Curve25519Field.A24) - val ddPlusTmp = Curve25519Field.add(dd, tmp) - val cNew = Curve25519Field.mul(cc, ddPlusTmp) - - aNew.copyInto(a) - bNew.copyInto(b) - cNew.copyInto(c) - dNew.copyInto(d) + Curve25519Field.mulInto(a, dd, ff, t) + Curve25519Field.subInto(cc, dd, ff) + Curve25519Field.mulInto(tmp, cc, Curve25519Field.A24, t) + Curve25519Field.addInto(tmp, dd, tmp) + Curve25519Field.mulInto(c, cc, tmp, t) Curve25519Field.sel25519(a, b, r) Curve25519Field.sel25519(c, d, r) } - val invC = Curve25519Field.inv25519(c) - val result = Curve25519Field.mul(a, invC) - return Curve25519Field.pack25519(result) + // c := 1/c, then a := a/c. `tmp` is free again and serves as the + // inversion's scratch element. + Curve25519Field.inv25519Into(c, c, tmp, t) + Curve25519Field.mulInto(a, a, c, t) + return Curve25519Field.pack25519(a) } } diff --git a/quartz/src/linuxMain/kotlin/com/vitorpamplona/quartz/marmot/mls/crypto/Ed25519.linux.kt b/quartz/src/linuxMain/kotlin/com/vitorpamplona/quartz/marmot/mls/crypto/Ed25519.linux.kt index 6339974e9d..1475607b2f 100644 --- a/quartz/src/linuxMain/kotlin/com/vitorpamplona/quartz/marmot/mls/crypto/Ed25519.linux.kt +++ b/quartz/src/linuxMain/kotlin/com/vitorpamplona/quartz/marmot/mls/crypto/Ed25519.linux.kt @@ -155,33 +155,68 @@ actual object Ed25519 { p[2].copyOf(), p[3].copyOf(), ) - addPointInPlace(result, q) + addPointInPlace(result, q, PointAddScratch()) return result } + /** + * Scratch for [addPointInPlace]. + * + * Extended-coordinate addition needs ten temporary field elements, and a + * scalar multiplication calls it 512 times — twice per scalar bit. Making + * each call allocate its own was the single largest source of garbage in + * the MLS stack: an allocation profile put 93% of all sampled allocation + * in `Curve25519Field.mul/add/sub`, most of it reached from here. + * + * One instance is created per scalar multiplication and reused by every + * step, so the loop allocates nothing. + */ + private class PointAddScratch { + val a = LongArray(16) + val b = LongArray(16) + val c = LongArray(16) + val d = LongArray(16) + val e = LongArray(16) + val f = LongArray(16) + val g = LongArray(16) + val h = LongArray(16) + val t1 = LongArray(16) + val t2 = LongArray(16) + + /** The 31-limb accumulator every [Curve25519Field.mulInto] here shares. */ + val mulT = LongArray(31) + } + private fun addPointInPlace( p: Array, q: Array, + s: PointAddScratch, ) { - val a = Curve25519Field.sub(p[1], p[0]) - val t = Curve25519Field.sub(q[1], q[0]) - val aMul = Curve25519Field.mul(a, t) - val b = Curve25519Field.add(p[0], p[1]) - val t2 = Curve25519Field.add(q[0], q[1]) - val bMul = Curve25519Field.mul(b, t2) - val c = Curve25519Field.mul(p[3], q[3]) - val cMul = Curve25519Field.mul(c, Curve25519Field.D2) - val d = Curve25519Field.mul(p[2], q[2]) - val dAdd = Curve25519Field.add(d, d) - val e = Curve25519Field.sub(bMul, aMul) - val f = Curve25519Field.sub(dAdd, cMul) - val g = Curve25519Field.add(dAdd, cMul) - val h = Curve25519Field.add(bMul, aMul) + // Every read of p and q happens in this first block. `scalarMult` + // doubles by passing the same point as both arguments, so nothing may + // be written back until these are done. + Curve25519Field.subInto(s.a, p[1], p[0]) + Curve25519Field.subInto(s.t1, q[1], q[0]) + Curve25519Field.mulInto(s.a, s.a, s.t1, s.mulT) + Curve25519Field.addInto(s.b, p[0], p[1]) + Curve25519Field.addInto(s.t2, q[0], q[1]) + Curve25519Field.mulInto(s.b, s.b, s.t2, s.mulT) + Curve25519Field.mulInto(s.c, p[3], q[3], s.mulT) + Curve25519Field.mulInto(s.c, s.c, Curve25519Field.D2, s.mulT) + Curve25519Field.mulInto(s.d, p[2], q[2], s.mulT) + Curve25519Field.addInto(s.d, s.d, s.d) - Curve25519Field.mul(e, f).copyInto(p[0]) - Curve25519Field.mul(h, g).copyInto(p[1]) - Curve25519Field.mul(g, f).copyInto(p[2]) - Curve25519Field.mul(e, h).copyInto(p[3]) + Curve25519Field.subInto(s.e, s.b, s.a) + Curve25519Field.subInto(s.f, s.d, s.c) + Curve25519Field.addInto(s.g, s.d, s.c) + Curve25519Field.addInto(s.h, s.b, s.a) + + // Safe to write p now: e, f, g and h are scratch, so no later product + // reads anything we are about to overwrite. + Curve25519Field.mulInto(p[0], s.e, s.f, s.mulT) + Curve25519Field.mulInto(p[1], s.h, s.g, s.mulT) + Curve25519Field.mulInto(p[2], s.g, s.f, s.mulT) + Curve25519Field.mulInto(p[3], s.e, s.h, s.mulT) } private fun negatePoint(p: Array): Array { @@ -205,11 +240,13 @@ actual object Ed25519 { p[2].copyOf(), p[3].copyOf(), ) + // One workspace for all 512 additions below. + val scratch = PointAddScratch() for (i in 255 downTo 0) { val b = ((s[i shr 3].toInt() shr (i and 7)) and 1).toLong() cswap(result, q, b) - addPointInPlace(q, result) - addPointInPlace(result, result) + addPointInPlace(q, result, scratch) + addPointInPlace(result, result, scratch) cswap(result, q, b) } return result diff --git a/quartz/src/linuxMain/kotlin/com/vitorpamplona/quartz/marmot/mls/crypto/X25519.linux.kt b/quartz/src/linuxMain/kotlin/com/vitorpamplona/quartz/marmot/mls/crypto/X25519.linux.kt index 90d78f9955..ed0d6a7ffa 100644 --- a/quartz/src/linuxMain/kotlin/com/vitorpamplona/quartz/marmot/mls/crypto/X25519.linux.kt +++ b/quartz/src/linuxMain/kotlin/com/vitorpamplona/quartz/marmot/mls/crypto/X25519.linux.kt @@ -68,6 +68,7 @@ actual object X25519 { p: ByteArray, ): ByteArray { val z = n.copyOf() + // Clamp scalar per RFC 7748 Section 5 z[0] = (z[0].toInt() and 248).toByte() z[31] = ((z[31].toInt() and 127) or 64).toByte() @@ -77,45 +78,62 @@ actual object X25519 { val c = Curve25519Field.GF0.copyOf() val d = Curve25519Field.GF1.copyOf() + // The ladder's entire working set, allocated ONCE. Every field + // operation in the loop writes into one of these, so 255 iterations + // allocate nothing at all — where the allocating form produced a fresh + // element per operation, about 1.3 MB of garbage per call. + val e = LongArray(16) + val f = LongArray(16) + val g = LongArray(16) + val h = LongArray(16) + val dd = LongArray(16) + val ff = LongArray(16) + val da = LongArray(16) + val cb = LongArray(16) + val cc = LongArray(16) + val tmp = LongArray(16) + val t = LongArray(31) + for (i in 254 downTo 0) { val r = ((z[i shr 3].toLong() shr (i and 7)) and 1) Curve25519Field.sel25519(a, b, r) Curve25519Field.sel25519(c, d, r) - val e = Curve25519Field.add(a, c) - val aMc = Curve25519Field.sub(a, c) - val f = Curve25519Field.add(b, d) - val bMd = Curve25519Field.sub(b, d) + // a, b, c and d are read only by these four lines; from here on + // they are dead and can be overwritten with the new values. + Curve25519Field.addInto(e, a, c) + Curve25519Field.subInto(g, a, c) + Curve25519Field.addInto(f, b, d) + Curve25519Field.subInto(h, b, d) - val dd = Curve25519Field.sqr(e) - val ff = Curve25519Field.sqr(aMc) - val da = Curve25519Field.mul(bMd, e) - val cb = Curve25519Field.mul(f, aMc) + Curve25519Field.sqrInto(dd, e, t) + Curve25519Field.sqrInto(ff, g, t) + Curve25519Field.mulInto(da, h, e, t) + Curve25519Field.mulInto(cb, f, g, t) - val ePrime = Curve25519Field.add(da, cb) - val aPrime = Curve25519Field.sub(da, cb) + // e := da + cb and g := da - cb. Reusing e and g is safe: both + // held inputs to the four products above, which are now computed. + Curve25519Field.addInto(e, da, cb) + Curve25519Field.subInto(g, da, cb) - val bNew = Curve25519Field.sqr(ePrime) - val aSqr = Curve25519Field.sqr(aPrime) - val dNew = Curve25519Field.mul(aSqr, x) + Curve25519Field.sqrInto(b, e, t) + Curve25519Field.sqrInto(g, g, t) + Curve25519Field.mulInto(d, g, x, t) - val aNew = Curve25519Field.mul(dd, ff) - val cc = Curve25519Field.sub(dd, ff) - val tmp = Curve25519Field.mul(cc, Curve25519Field.A24) - val ddPlusTmp = Curve25519Field.add(dd, tmp) - val cNew = Curve25519Field.mul(cc, ddPlusTmp) - - aNew.copyInto(a) - bNew.copyInto(b) - cNew.copyInto(c) - dNew.copyInto(d) + Curve25519Field.mulInto(a, dd, ff, t) + Curve25519Field.subInto(cc, dd, ff) + Curve25519Field.mulInto(tmp, cc, Curve25519Field.A24, t) + Curve25519Field.addInto(tmp, dd, tmp) + Curve25519Field.mulInto(c, cc, tmp, t) Curve25519Field.sel25519(a, b, r) Curve25519Field.sel25519(c, d, r) } - val invC = Curve25519Field.inv25519(c) - val result = Curve25519Field.mul(a, invC) - return Curve25519Field.pack25519(result) + // c := 1/c, then a := a/c. `tmp` is free again and serves as the + // inversion's scratch element. + Curve25519Field.inv25519Into(c, c, tmp, t) + Curve25519Field.mulInto(a, a, c, t) + return Curve25519Field.pack25519(a) } }