From 58e684273eff0f647cc28483029b5c4d197c98e4 Mon Sep 17 00:00:00 2001 From: Claude Date: Thu, 10 Sep 2026 13:47:11 +0000 Subject: [PATCH] perf(marmot): run Curve25519 scalar multiplication without allocating An allocation profile of the Marmot benchmarks put 93% of every sampled allocation in Curve25519Field.mul/add/sub. The pure-Kotlin field arithmetic returned a fresh LongArray(16) from every operation, and a Montgomery ladder runs ~18 of them per scalar bit across 255 bits, so one X25519 scalar multiplication produced over a megabyte of garbage. Ed25519 was worse: its extended-coordinate point addition needs ten temporaries and a scalar multiplication calls it 512 times. Give each field operation an *Into twin that writes into a caller-owned output and shares one 31-limb accumulator, then rewrite both hot paths around them. The X25519 ladder allocates its eleven-array working set once before the loop and overwrites a/b/c/d in place after their last read; Ed25519 creates a single PointAddScratch per scalar multiplication and reuses it for all 512 additions, including the aliasing doubling step. Every *Into is safe when the output aliases an input, because mulInto fully accumulates into the scratch before it touches the output. The allocating functions stay. They are still used off the hot path, where clarity is worth more than the bytes, and keeping them means the in-place versions can be differentially tested against them. Allocation per operation drops 10x to 80x depending on the benchmark (create_group/0 6958.7 KB to 88.3 KB, ingest_app_message 5640.3 KB to 70.9 KB), reproducing to four significant figures across runs, and p50 latency improves on every row that is not dominated by measurement noise. No behaviour changes: the RFC 7748 and RFC 8032 vector suites, the HPKE tests and the full quartz suite pass unchanged. Co-Authored-By: Claude Opus 5 Claude-Session: https://claude.ai/code/session_016kCuA6tc4JQzHPCDd39GHq --- marmotBench/README.md | 48 ++++++++ .../quartz/marmot/mls/crypto/Ed25519.apple.kt | 79 +++++++++---- .../quartz/marmot/mls/crypto/X25519.apple.kt | 69 ++++++----- .../marmot/mls/crypto/Curve25519Field.kt | 108 +++++++++++++++--- .../marmot/mls/crypto/Ed25519.jvmAndroid.kt | 79 +++++++++---- .../marmot/mls/crypto/X25519.jvmAndroid.kt | 69 ++++++----- .../quartz/marmot/mls/crypto/Ed25519.linux.kt | 79 +++++++++---- .../quartz/marmot/mls/crypto/X25519.linux.kt | 70 +++++++----- 8 files changed, 447 insertions(+), 154 deletions(-) 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) } }