diff --git a/marmotBench/README.md b/marmotBench/README.md index 0d151b2d57..c71caed009 100644 --- a/marmotBench/README.md +++ b/marmotBench/README.md @@ -108,3 +108,64 @@ 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. + +## Result: 10 limbs instead of 16 + +The allocation work above left `create_group` still ~3.9x slower than MDK, and +the primitive benchmarks said why: one X25519 scalar multiplication cost 541us, +and `create_group/1` is about two dozen of them. Curve work *was* the operation. + +The cause was the representation, not the language. `Curve25519Field` used +TweetNaCl's 16 limbs of radix 2^16, so a schoolbook field multiply spent 256 +limb products. SunEC's X25519 — also pure Java, same JIT, same machine — ran +the same operation in 160us using ~26-bit limbs in 10 words, which is 100 +products. The ratio of products matched the ratio of times. + +So the field was rewritten to 10 limbs of radix 2^25.5, the layout ref10, +curve25519-donna and SunEC all use: 100 products per multiply, 55 per square +(each off-diagonal pair once, doubled), and a dedicated scalar multiply for the +ladder's a24 constant instead of a general multiply against nine zero limbs. + +| primitive | 16 limbs | 10 limbs | speedup | +|------------------|----------|----------|---------| +| `x25519_dh` | 541us | 121us | 4.5x | +| `x25519_base` | 535us | 121us | 4.4x | +| `ed25519_sign` | 1018us | 259us | 3.9x | +| `ed25519_verify` | 2113us | 536us | 3.9x | + +At 121us the scalar multiplication is now faster than SunEC's 160us, which is +the useful sanity check on the result: it lands where a good managed-language +implementation should, rather than somewhere suspiciously better. + +Against MDK, over the whole suite (both post-rewrite runs shown where they +differ; `alloc/op` reproduces to four significant figures): + +| operation | MDK (Rust) | quartz before | quartz now | vs MDK | +|----------------------|------------|---------------|-----------------|--------------| +| `create_group/1` | 3.61 ms | 16.93 ms | 5.42 - 5.88 ms | 1.5-1.6x slower | +| `create_group/8` | 9.93 ms | 51.47 ms | 17.27 - 17.60 ms| 1.8x slower | +| `create_group/32` | 31.64 ms | 190.08 ms | 77.53 - 81.47 ms| 2.5x slower | +| `join_welcome` | 4.77 ms | 6.22 ms | 1.82 - 2.00 ms | **2.5x faster** | +| `send_app_message` | 4.28 ms | 1.72 ms | 0.61 - 0.71 ms | **6.5x faster** | +| `ingest_app_message` | (n/a) | 3.11 ms | 0.88 - 0.90 ms | — | + +`create_group` remains the weakest row, and the shape difference in "What is +compared" is part of why: we create at epoch 0 and add in a second commit, +where MDK folds invitees into the founding group. `create_group/32` is also +still the noisiest row in the suite. + +### Why the constants can be trusted + +Changing the representation re-encodes every curve constant, which is exactly +the kind of change where a single mistyped limb produces code that still runs +and is still wrong. None of them were transcribed by hand: each was re-derived +from its existing 16-bit encoding and then checked against its mathematical +definition — `d == -121665/121666`, `d2 == 2d`, `By == 4/5`, `I^2 == -1` — and +the multiply and square formulas were generated from the representation's +weight bookkeeping and diffed against an independent reference over 20 000 +random limb vectors before any Kotlin was written. The carry chain and the +canonical encoder were validated the same way, including at `p`, `p-1`, and on +non-canonical inputs such as `p` itself. + +The RFC 7748 and RFC 8032 vector suites, HPKE, the MDK crypto-interop vectors +and the full quartz + commons suites all pass unchanged. 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 6f6d950d46..74c7621bc8 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 @@ -148,15 +148,15 @@ actual object Ed25519 { } // --- Extended Edwards point operations --- - // Point = Array of 4 LongArray(16), representing (X, Y, Z, T) + // Point = Array of 4 LongArray(10), representing (X, Y, Z, T) // where x = X/Z, y = Y/Z, x*y = T/Z private fun newPoint(): Array = arrayOf( - LongArray(16), - LongArray(16), - LongArray(16), - LongArray(16), + LongArray(10), + LongArray(10), + LongArray(10), + LongArray(10), ) /** Set point to the identity (0, 1, 1, 0). */ @@ -196,19 +196,16 @@ actual object Ed25519 { * 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) + val a = LongArray(10) + val b = LongArray(10) + val c = LongArray(10) + val d = LongArray(10) + val e = LongArray(10) + val f = LongArray(10) + val g = LongArray(10) + val h = LongArray(10) + val t1 = LongArray(10) + val t2 = LongArray(10) } /** In-place point addition: p += q. */ @@ -222,13 +219,13 @@ actual object Ed25519 { // 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.mulInto(s.a, s.a, s.t1) 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.mulInto(s.b, s.b, s.t2) + Curve25519Field.mulInto(s.c, p[3], q[3]) + Curve25519Field.mulInto(s.c, s.c, Curve25519Field.D2) + Curve25519Field.mulInto(s.d, p[2], q[2]) Curve25519Field.addInto(s.d, s.d, s.d) Curve25519Field.subInto(s.e, s.b, s.a) @@ -238,10 +235,10 @@ actual object Ed25519 { // 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) + Curve25519Field.mulInto(p[0], s.e, s.f) + Curve25519Field.mulInto(p[1], s.h, s.g) + Curve25519Field.mulInto(p[2], s.g, s.f) + Curve25519Field.mulInto(p[3], s.e, s.h) } /** Point doubling (self-addition). */ @@ -325,25 +322,7 @@ actual object Ed25519 { // Recover x from y: x^2 = (y^2 - 1) / (d * y^2 + 1) val y2 = Curve25519Field.sqr(r) - val d = - Curve25519Field.gf( - 0x78A3, - 0x1359, - 0x4DCA, - 0x75EB, - 0xD8AB, - 0x4141, - 0x0A4D, - 0x0070, - 0xE898, - 0x7779, - 0x4079, - 0x8CC7, - 0xFE73, - 0x2B6F, - 0x6CEE, - 0x5203, - ) + val d = Curve25519Field.D val num = Curve25519Field.sub(y2, Curve25519Field.GF1) val den = Curve25519Field.add(Curve25519Field.mul(d, y2), Curve25519Field.GF1) val denInv = Curve25519Field.inv25519(den) 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 1625e6b4f4..6577f257ed 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 @@ -85,17 +85,16 @@ actual object X25519 { // 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) + val e = LongArray(10) + val f = LongArray(10) + val g = LongArray(10) + val h = LongArray(10) + val dd = LongArray(10) + val ff = LongArray(10) + val da = LongArray(10) + val cb = LongArray(10) + val cc = LongArray(10) + val tmp = LongArray(10) for (i in 254 downTo 0) { val r = ((z[i shr 3].toLong() shr (i and 7)) and 1) @@ -109,25 +108,25 @@ actual object X25519 { Curve25519Field.addInto(f, b, d) Curve25519Field.subInto(h, b, d) - Curve25519Field.sqrInto(dd, e, t) - Curve25519Field.sqrInto(ff, g, t) - Curve25519Field.mulInto(da, h, e, t) - Curve25519Field.mulInto(cb, f, g, t) + Curve25519Field.sqrInto(dd, e) + Curve25519Field.sqrInto(ff, g) + Curve25519Field.mulInto(da, h, e) + Curve25519Field.mulInto(cb, f, g) // 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) - Curve25519Field.sqrInto(b, e, t) - Curve25519Field.sqrInto(g, g, t) - Curve25519Field.mulInto(d, g, x, t) + Curve25519Field.sqrInto(b, e) + Curve25519Field.sqrInto(g, g) + Curve25519Field.mulInto(d, g, x) - Curve25519Field.mulInto(a, dd, ff, t) + Curve25519Field.mulInto(a, dd, ff) Curve25519Field.subInto(cc, dd, ff) - Curve25519Field.mulInto(tmp, cc, Curve25519Field.A24, t) + Curve25519Field.mulA24Into(tmp, cc) Curve25519Field.addInto(tmp, dd, tmp) - Curve25519Field.mulInto(c, cc, tmp, t) + Curve25519Field.mulInto(c, cc, tmp) Curve25519Field.sel25519(a, b, r) Curve25519Field.sel25519(c, d, r) @@ -135,8 +134,8 @@ actual object X25519 { // 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) + Curve25519Field.inv25519Into(c, c, tmp) + Curve25519Field.mulInto(a, a, c) 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 2a6ed3cc14..97ee380680 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 @@ -21,238 +21,93 @@ package com.vitorpamplona.quartz.marmot.mls.crypto /** - * Field arithmetic over GF(2^255-19) for Curve25519 operations. + * Field arithmetic over GF(2^255-19) for Curve25519 and Ed25519. * - * Field elements are represented as LongArray(16) in radix-2^16. - * Based on the TweetNaCl algorithm by Bernstein et al. + * ## Representation: 10 limbs, radix 2^25.5 + * + * A field element is a `LongArray(10)`. Limb `i` carries weight `2^OFFSET[i]` + * where the offsets step alternately by 26 and 25 bits — even limbs hold 26 + * bits, odd limbs 25 — so ten limbs span the 255 bits of the field. This is + * the layout ref10, curve25519-donna and SunEC all use. + * + * It replaces TweetNaCl's 16 limbs of radix 2^16, and the reason is arithmetic + * rather than taste. A schoolbook multiply costs one product per pair of + * limbs: 16 limbs means 256 products, 10 limbs means 100. Measured on the + * benchmark machine, that is the difference between a 541us scalar + * multiplication and SunEC's 160us — and SunEC is itself pure Java on the same + * JIT, which is what rules out "the JVM is slow" as the explanation. + * + * Limbs are SIGNED. Subtraction does not borrow and multiplication does not + * normalise beyond a carry chain, so intermediate limbs are allowed to go + * negative and to exceed their nominal width; only [pack25519] produces a + * canonical value. The bound that matters is that a limb entering [mulInto] + * stays under about 2^26 in absolute value, which leaves the widest + * accumulator column at 2^60 — three bits clear of overflowing a signed 64-bit + * Long. Every operation here preserves that. + * + * ## Allocating vs in-place + * + * Each `add`/`sub`/`mul`/`sqr` has an `*Into` twin that writes into a + * caller-owned output, because the allocating forms turned a scalar + * multiplication into about a megabyte of garbage. The hot paths (the X25519 + * ladder and Ed25519 point addition) use the in-place forms exclusively and + * allocate nothing; the allocating forms remain for off-hot-path clarity and + * as a differential-testing partner for the in-place ones. + * + * Unlike the 16-limb version these need no scratch accumulator: [mulInto] and + * [sqrInto] are straight-line over local Longs, so there is no array to pass + * in and none to zero. + * + * Every `*Into` is safe when the output aliases an input — the ladder relies + * on that, and it holds because each reads every input into locals before + * writing any output. */ internal object Curve25519Field { - /** The constant a24 = 121665, used in the Montgomery ladder. */ - val A24 = gf(0xDB41L, 1) + /** Bit offset of each limb: even limbs are 26 bits wide, odd limbs 25. */ + private val OFFSET = intArrayOf(0, 26, 51, 77, 102, 128, 153, 179, 204, 230) - /** d2 = 2*d where d is the Edwards curve constant, for point addition. */ - val D2 = - gf( - 0xF159, - 0x26B2, - 0x9B94, - 0xEBD6, - 0xB156, - 0x8283, - 0x149A, - 0x00E0, - 0xD130, - 0xEEF3, - 0x80F2, - 0x198E, - 0xFCE7, - 0x56DF, - 0xD9DC, - 0x2406, - ) - - /** Ed25519 base point X coordinate. */ - val BX = - gf( - 0xD51A, - 0x8F25, - 0x2D60, - 0xC956, - 0xA7B2, - 0x9525, - 0xC760, - 0x692C, - 0xDC5C, - 0xFDD6, - 0xE231, - 0xC0A4, - 0x53FE, - 0xCD6E, - 0x36D3, - 0x2169, - ) - - /** Ed25519 base point Y coordinate. */ - val BY = - gf( - 0x6658, - 0x6666, - 0x6666, - 0x6666, - 0x6666, - 0x6666, - 0x6666, - 0x6666, - 0x6666, - 0x6666, - 0x6666, - 0x6666, - 0x6666, - 0x6666, - 0x6666, - 0x6666, - ) - - /** sqrt(-1) mod p, used for Ed25519 point decompression. */ - val I = - gf( - 0xA0B0, - 0x4A0E, - 0x1B27, - 0xC4EE, - 0xE478, - 0xAD2F, - 0x1806, - 0x2F43, - 0xD7A7, - 0x3DFB, - 0x0099, - 0x2B4D, - 0xDF0B, - 0x4FC1, - 0x2480, - 0x2B83, - ) - - fun gf(vararg values: Long): LongArray { - val o = LongArray(16) - for (i in values.indices) { - o[i] = values[i] - } - return o - } - - fun gf( - a: Long, - b: Long, - ): LongArray { - val o = LongArray(16) - o[0] = a - o[1] = b - return o - } - - val GF0 = LongArray(16) - val GF1 = gf(1) + /** Number of limbs in a field element. */ + const val LIMBS = 10 /** - * Carry and reduce a field element. + * Build a field element from its limbs, zero-filling the rest. * - * TweetNaCl writes the loop over all 16 limbs and folds the wrap-around - * into the body as `o[(i + 1) % 16]` plus an `if (i == 15)`, so a modulo - * and a branch ride along on all 16 iterations to serve the one that needs - * them. Peeling the last limb out takes both off the loop: limbs 0..14 - * carry into their neighbour, and limb 15 wraps into limb 0 scaled by 38, - * which is the `c - 1` plus the `37 * (c - 1)` of the original folded into - * one term. - * - * Measured honestly, this bought **nothing** on HotSpot — C2 was already - * strength-reducing the modulo and hoisting the branch. It is kept because - * it is strictly less work for a weaker JIT to undo, and ART on a phone is - * the target that matters, but no speedup is claimed for it here: the - * benchmark on this machine could not tell the two apart. - * - * That measurement is also the reason not to trust a CPU profile of this - * file. JFR's execution sampler is safepoint-biased, and the counted loops - * in this object carry no safepoint polls, so samples pile onto whichever - * method follows the poll. It attributed 75% of all `create_group` samples - * to this function; rewriting it changed nothing, which is the profiler - * telling on itself. Time the primitives end to end instead — see - * `marmotBench`'s `x25519_dh` and friends. - * - * The `+ (1 shl 16)` / `- 1` dance is TweetNaCl's, and stays: it biases the - * limb so an arithmetic shift floors correctly for negative limbs, which is - * what makes the carry branch-free for the sign as well. + * Values are limbs in THIS representation, not bytes — see [unpack25519] + * to go from a 32-byte encoding. */ - fun car25519(o: LongArray) { - for (i in 0 until 15) { - o[i] += (1L shl 16) - val c = o[i] shr 16 - o[i + 1] += c - 1 - o[i] -= c shl 16 - } - o[15] += (1L shl 16) - val c = o[15] shr 16 - o[0] += 38 * (c - 1) - o[15] -= c shl 16 - } - - /** Conditional swap: if b=1, swap p and q element-wise. */ - fun sel25519( - p: LongArray, - q: LongArray, - b: Long, - ) { - val c = b.inv() + 1 // 0 -> 0, 1 -> -1 (all ones) - for (i in 0 until 16) { - val t = c and (p[i] xor q[i]) - p[i] = p[i] xor t - q[i] = q[i] xor t - } - } - - /** Pack a field element to 32-byte little-endian representation. */ - fun pack25519(n: LongArray): ByteArray { - val o = ByteArray(32) - val m = LongArray(16) - val t = n.copyOf() - car25519(t) - car25519(t) - car25519(t) - for (j in 0 until 2) { - m[0] = t[0] - 0xFFED - for (i in 1 until 15) { - m[i] = t[i] - 0xFFFF - ((m[i - 1] shr 16) and 1) - m[i - 1] = m[i - 1] and 0xFFFF - } - m[15] = t[15] - 0x7FFF - ((m[14] shr 16) and 1) - val b = (m[15] shr 16) and 1 - m[14] = m[14] and 0xFFFF - sel25519(t, m, 1 - b) - } - for (i in 0 until 16) { - o[2 * i] = (t[i] and 0xFF).toByte() - o[2 * i + 1] = (t[i] shr 8).toByte() - } + fun gf(vararg values: Long): LongArray { + val o = LongArray(LIMBS) + for (i in values.indices) o[i] = values[i] return o } - /** Unpack 32-byte little-endian to field element. */ - fun unpack25519(n: ByteArray): LongArray { - val o = LongArray(16) - for (i in 0 until 16) { - o[i] = (n[2 * i].toLong() and 0xFF) + ((n[2 * i + 1].toLong() and 0xFF) shl 8) - } - o[15] = o[15] and 0x7FFF - return o - } + val GF0 = LongArray(LIMBS) + val GF1 = gf(1) - // 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. + /** a24 = 121665, the Montgomery ladder constant. */ + val A24 = gf(121665) + + /** d = -121665/121666, the Edwards curve constant. */ + val D = gf(56195235, 13857412, 51736253, 6949390, 114729, 24766616, 60832955, 30306712, 48412415, 21499315) + + /** d2 = 2*d, for extended-coordinate point addition. */ + val D2 = gf(45281625, 27714825, 36363642, 13898781, 229458, 15978800, 54557047, 27058993, 29715967, 9444199) + + /** Ed25519 base point X coordinate. */ + val BX = gf(52811034, 25909283, 16144682, 17082669, 27570973, 30858332, 40966398, 8378388, 20764389, 8758491) + + /** Ed25519 base point Y coordinate, which is 4/5. */ + val BY = gf(40265304, 26843545, 13421772, 20132659, 26843545, 6710886, 53687091, 13421772, 40265318, 26843545) + + /** sqrt(-1) mod p, used for Ed25519 point decompression. */ + val I = gf(34513072, 25610706, 9377949, 3500415, 12389472, 33281959, 41962654, 31548777, 326685, 11406482) /** Field addition: o = a + b. */ fun add( a: LongArray, b: LongArray, ): LongArray { - val o = LongArray(16) + val o = LongArray(LIMBS) addInto(o, a, b) return o } @@ -263,7 +118,7 @@ internal object Curve25519Field { a: LongArray, b: LongArray, ) { - for (i in 0 until 16) o[i] = a[i] + b[i] + for (i in 0 until LIMBS) o[i] = a[i] + b[i] } /** Field subtraction: o = a - b. */ @@ -271,7 +126,7 @@ internal object Curve25519Field { a: LongArray, b: LongArray, ): LongArray { - val o = LongArray(16) + val o = LongArray(LIMBS) subInto(o, a, b) return o } @@ -282,7 +137,7 @@ internal object Curve25519Field { a: LongArray, b: LongArray, ) { - for (i in 0 until 16) o[i] = a[i] - b[i] + for (i in 0 until LIMBS) o[i] = a[i] - b[i] } /** Field multiplication: o = a * b (mod p). */ @@ -290,111 +145,431 @@ internal object Curve25519Field { a: LongArray, b: LongArray, ): LongArray { - val o = LongArray(16) - mulInto(o, a, b, LongArray(31)) + val o = LongArray(LIMBS) + mulInto(o, a, b) return o } /** - * Field multiplication into [o], using [t] as the 31-limb accumulator. + * Field multiplication into [o]: o = f * g (mod p). * - * [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. + * Straight-line over locals, so it allocates nothing and [o] may alias + * either input. Generated from the weight bookkeeping of the + * representation and cross-checked against an independent reference: a + * product of limbs i and j lands in limb i+j, doubled when i and j are + * both odd (their offsets sum to one more than the target's), and folded + * back into limb i+j-10 scaled by 19 when it overflows the top, since + * 2^255 == 19 (mod p). + * + * 100 products, against 256 for the 16-limb representation this replaced. */ fun mulInto( o: LongArray, - a: LongArray, - b: LongArray, - t: LongArray, + f: LongArray, + g: LongArray, ) { - t.fill(0L) - for (i in 0 until 16) { - val ai = a[i] - for (j in 0 until 16) { - t[i + j] += ai * b[j] - } - } - for (i in 0 until 15) { - t[i] += 38 * t[i + 16] - } - for (i in 0 until 16) o[i] = t[i] - car25519(o) - car25519(o) + val f0 = f[0] + val f1 = f[1] + val f2 = f[2] + val f3 = f[3] + val f4 = f[4] + val f5 = f[5] + val f6 = f[6] + val f7 = f[7] + val f8 = f[8] + val f9 = f[9] + val g0 = g[0] + val g1 = g[1] + val g2 = g[2] + val g3 = g[3] + val g4 = g[4] + val g5 = g[5] + val g6 = g[6] + val g7 = g[7] + val g8 = g[8] + val g9 = g[9] + val f1x2 = f1 + f1 + val f3x2 = f3 + f3 + val f5x2 = f5 + f5 + val f7x2 = f7 + f7 + val f9x2 = f9 + f9 + val g1x19 = 19 * g1 + val g2x19 = 19 * g2 + val g3x19 = 19 * g3 + val g4x19 = 19 * g4 + val g5x19 = 19 * g5 + val g6x19 = 19 * g6 + val g7x19 = 19 * g7 + val g8x19 = 19 * g8 + val g9x19 = 19 * g9 + var h0 = f0 * g0 + f1x2 * g9x19 + f2 * g8x19 + f3x2 * g7x19 + f4 * g6x19 + f5x2 * g5x19 + f6 * g4x19 + f7x2 * g3x19 + f8 * g2x19 + f9x2 * g1x19 + var h1 = f0 * g1 + f1 * g0 + f2 * g9x19 + f3 * g8x19 + f4 * g7x19 + f5 * g6x19 + f6 * g5x19 + f7 * g4x19 + f8 * g3x19 + f9 * g2x19 + var h2 = f0 * g2 + f1x2 * g1 + f2 * g0 + f3x2 * g9x19 + f4 * g8x19 + f5x2 * g7x19 + f6 * g6x19 + f7x2 * g5x19 + f8 * g4x19 + f9x2 * g3x19 + var h3 = f0 * g3 + f1 * g2 + f2 * g1 + f3 * g0 + f4 * g9x19 + f5 * g8x19 + f6 * g7x19 + f7 * g6x19 + f8 * g5x19 + f9 * g4x19 + var h4 = f0 * g4 + f1x2 * g3 + f2 * g2 + f3x2 * g1 + f4 * g0 + f5x2 * g9x19 + f6 * g8x19 + f7x2 * g7x19 + f8 * g6x19 + f9x2 * g5x19 + var h5 = f0 * g5 + f1 * g4 + f2 * g3 + f3 * g2 + f4 * g1 + f5 * g0 + f6 * g9x19 + f7 * g8x19 + f8 * g7x19 + f9 * g6x19 + var h6 = f0 * g6 + f1x2 * g5 + f2 * g4 + f3x2 * g3 + f4 * g2 + f5x2 * g1 + f6 * g0 + f7x2 * g9x19 + f8 * g8x19 + f9x2 * g7x19 + var h7 = f0 * g7 + f1 * g6 + f2 * g5 + f3 * g4 + f4 * g3 + f5 * g2 + f6 * g1 + f7 * g0 + f8 * g9x19 + f9 * g8x19 + var h8 = f0 * g8 + f1x2 * g7 + f2 * g6 + f3x2 * g5 + f4 * g4 + f5x2 * g3 + f6 * g2 + f7x2 * g1 + f8 * g0 + f9x2 * g9x19 + var h9 = f0 * g9 + f1 * g8 + f2 * g7 + f3 * g6 + f4 * g5 + f5 * g4 + f6 * g3 + f7 * g2 + f8 * g1 + f9 * g0 + // ref10's carry ordering: independent carries are interleaved so the + // limb-to-limb dependency chain does not stall the pipeline. + var c = (h0 + (1L shl 25)) shr 26 + h1 += c + h0 -= c shl 26 + c = (h4 + (1L shl 25)) shr 26 + h5 += c + h4 -= c shl 26 + c = (h1 + (1L shl 24)) shr 25 + h2 += c + h1 -= c shl 25 + c = (h5 + (1L shl 24)) shr 25 + h6 += c + h5 -= c shl 25 + c = (h2 + (1L shl 25)) shr 26 + h3 += c + h2 -= c shl 26 + c = (h6 + (1L shl 25)) shr 26 + h7 += c + h6 -= c shl 26 + c = (h3 + (1L shl 24)) shr 25 + h4 += c + h3 -= c shl 25 + c = (h7 + (1L shl 24)) shr 25 + h8 += c + h7 -= c shl 25 + c = (h4 + (1L shl 25)) shr 26 + h5 += c + h4 -= c shl 26 + c = (h8 + (1L shl 25)) shr 26 + h9 += c + h8 -= c shl 26 + c = (h9 + (1L shl 24)) shr 25 + h0 += 19 * c + h9 -= c shl 25 + c = (h0 + (1L shl 25)) shr 26 + h1 += c + h0 -= c shl 26 + o[0] = h0 + o[1] = h1 + o[2] = h2 + o[3] = h3 + o[4] = h4 + o[5] = h5 + o[6] = h6 + o[7] = h7 + o[8] = h8 + o[9] = h9 } /** Field squaring: o = a^2 (mod p). */ fun sqr(a: LongArray): LongArray { - val o = LongArray(16) - sqrInto(o, a, LongArray(31)) + val o = LongArray(LIMBS) + sqrInto(o, a) return o } /** - * Field squaring into [o]. See [mulInto] for the [t] contract. + * Field squaring into [o]: o = f^2 (mod p). * - * A square is not just `mulInto(o, a, a, t)`: in `a[i] * a[j]` every - * off-diagonal pair is computed twice, once as (i,j) and once as (j,i). - * Taking each pair once and doubling it turns the 256 multiplications of - * the schoolbook into 136 — the 16 diagonal squares plus 120 cross terms. - * - * That is worth having because squarings are not a rare case: the - * Montgomery ladder squares four times per bit out of ten field - * multiplications, and [inv25519Into] is 254 squarings against ~250 - * multiplications. - * - * Doubling costs no headroom. Limbs reaching here are bounded well under - * 2^18 even after an unreduced add or subtract, so a doubled cross term - * stays under 2^37 and a full 16-term column under 2^41 — far from - * overflowing the signed 64-bit accumulator. + * Same bookkeeping as [mulInto], with each off-diagonal pair taken once + * and doubled instead of computed twice: 55 products against the multiply's + * 100. Squarings are not a rare case — the ladder squares four times per + * bit, and [inv25519Into] is 254 squarings. */ fun sqrInto( o: LongArray, - a: LongArray, - t: LongArray, + f: LongArray, ) { - t.fill(0L) - for (i in 0 until 16) { - val ai = a[i] - t[i + i] += ai * ai - val twice = ai + ai - for (j in i + 1 until 16) { - t[i + j] += twice * a[j] + val f0 = f[0] + val f1 = f[1] + val f2 = f[2] + val f3 = f[3] + val f4 = f[4] + val f5 = f[5] + val f6 = f[6] + val f7 = f[7] + val f8 = f[8] + val f9 = f[9] + val f0x2 = f0 + f0 + val f1x2 = f1 + f1 + val f2x2 = f2 + f2 + val f3x2 = f3 + f3 + val f4x2 = f4 + f4 + val f5x2 = f5 + f5 + val f6x2 = f6 + f6 + val f7x2 = f7 + f7 + val f8x2 = f8 + f8 + val f9x2 = f9 + f9 + val f1x4 = 4 * f1 + val f3x4 = 4 * f3 + val f5x4 = 4 * f5 + val f7x4 = 4 * f7 + val f1x19 = 19 * f1 + val f2x19 = 19 * f2 + val f3x19 = 19 * f3 + val f4x19 = 19 * f4 + val f5x19 = 19 * f5 + val f6x19 = 19 * f6 + val f7x19 = 19 * f7 + val f8x19 = 19 * f8 + val f9x19 = 19 * f9 + var h0 = f0 * f0 + f1x4 * f9x19 + f2x2 * f8x19 + f3x4 * f7x19 + f4x2 * f6x19 + f5x2 * f5x19 + var h1 = f0x2 * f1 + f2x2 * f9x19 + f3x2 * f8x19 + f4x2 * f7x19 + f5x2 * f6x19 + var h2 = f0x2 * f2 + f1x2 * f1 + f3x4 * f9x19 + f4x2 * f8x19 + f5x4 * f7x19 + f6 * f6x19 + var h3 = f0x2 * f3 + f1x2 * f2 + f4x2 * f9x19 + f5x2 * f8x19 + f6x2 * f7x19 + var h4 = f0x2 * f4 + f1x4 * f3 + f2 * f2 + f5x4 * f9x19 + f6x2 * f8x19 + f7x2 * f7x19 + var h5 = f0x2 * f5 + f1x2 * f4 + f2x2 * f3 + f6x2 * f9x19 + f7x2 * f8x19 + var h6 = f0x2 * f6 + f1x4 * f5 + f2x2 * f4 + f3x2 * f3 + f7x4 * f9x19 + f8 * f8x19 + var h7 = f0x2 * f7 + f1x2 * f6 + f2x2 * f5 + f3x2 * f4 + f8x2 * f9x19 + var h8 = f0x2 * f8 + f1x4 * f7 + f2x2 * f6 + f3x4 * f5 + f4 * f4 + f9x2 * f9x19 + var h9 = f0x2 * f9 + f1x2 * f8 + f2x2 * f7 + f3x2 * f6 + f4x2 * f5 + // ref10's carry ordering: independent carries are interleaved so the + // limb-to-limb dependency chain does not stall the pipeline. + var c = (h0 + (1L shl 25)) shr 26 + h1 += c + h0 -= c shl 26 + c = (h4 + (1L shl 25)) shr 26 + h5 += c + h4 -= c shl 26 + c = (h1 + (1L shl 24)) shr 25 + h2 += c + h1 -= c shl 25 + c = (h5 + (1L shl 24)) shr 25 + h6 += c + h5 -= c shl 25 + c = (h2 + (1L shl 25)) shr 26 + h3 += c + h2 -= c shl 26 + c = (h6 + (1L shl 25)) shr 26 + h7 += c + h6 -= c shl 26 + c = (h3 + (1L shl 24)) shr 25 + h4 += c + h3 -= c shl 25 + c = (h7 + (1L shl 24)) shr 25 + h8 += c + h7 -= c shl 25 + c = (h4 + (1L shl 25)) shr 26 + h5 += c + h4 -= c shl 26 + c = (h8 + (1L shl 25)) shr 26 + h9 += c + h8 -= c shl 26 + c = (h9 + (1L shl 24)) shr 25 + h0 += 19 * c + h9 -= c shl 25 + c = (h0 + (1L shl 25)) shr 26 + h1 += c + h0 -= c shl 26 + o[0] = h0 + o[1] = h1 + o[2] = h2 + o[3] = h3 + o[4] = h4 + o[5] = h5 + o[6] = h6 + o[7] = h7 + o[8] = h8 + o[9] = h9 + } + + /** + * Multiply by the ladder constant a24 = 121665. + * + * One limb-wise scalar multiply plus a carry, instead of the 100 products + * a general [mulInto] would spend against a constant that is zero in nine + * of its ten limbs. The ladder does this once per bit. + * + * 121665 < 2^17 and limbs stay under 2^26, so each product stays under + * 2^43 — nowhere near overflowing. + */ + fun mulA24Into( + o: LongArray, + f: LongArray, + ) { + var h0 = f[0] * 121665 + var h1 = f[1] * 121665 + var h2 = f[2] * 121665 + var h3 = f[3] * 121665 + var h4 = f[4] * 121665 + var h5 = f[5] * 121665 + var h6 = f[6] * 121665 + var h7 = f[7] * 121665 + var h8 = f[8] * 121665 + var h9 = f[9] * 121665 + // ref10's carry ordering: independent carries are interleaved so the + // limb-to-limb dependency chain does not stall the pipeline. + var c = (h0 + (1L shl 25)) shr 26 + h1 += c + h0 -= c shl 26 + c = (h4 + (1L shl 25)) shr 26 + h5 += c + h4 -= c shl 26 + c = (h1 + (1L shl 24)) shr 25 + h2 += c + h1 -= c shl 25 + c = (h5 + (1L shl 24)) shr 25 + h6 += c + h5 -= c shl 25 + c = (h2 + (1L shl 25)) shr 26 + h3 += c + h2 -= c shl 26 + c = (h6 + (1L shl 25)) shr 26 + h7 += c + h6 -= c shl 26 + c = (h3 + (1L shl 24)) shr 25 + h4 += c + h3 -= c shl 25 + c = (h7 + (1L shl 24)) shr 25 + h8 += c + h7 -= c shl 25 + c = (h4 + (1L shl 25)) shr 26 + h5 += c + h4 -= c shl 26 + c = (h8 + (1L shl 25)) shr 26 + h9 += c + h8 -= c shl 26 + c = (h9 + (1L shl 24)) shr 25 + h0 += 19 * c + h9 -= c shl 25 + c = (h0 + (1L shl 25)) shr 26 + h1 += c + h0 -= c shl 26 + o[0] = h0 + o[1] = h1 + o[2] = h2 + o[3] = h3 + o[4] = h4 + o[5] = h5 + o[6] = h6 + o[7] = h7 + o[8] = h8 + o[9] = h9 + } + + /** + * Carry and partially reduce a field element in place. + * + * Brings limbs back inside their nominal widths so a value that has been + * added or subtracted repeatedly is safe to feed to [mulInto] again. It + * does NOT produce a canonical representative — [pack25519] does that. + */ + fun car25519(o: LongArray) { + var c: Long + for (i in 0 until LIMBS) { + val width = if (i and 1 == 0) 26 else 25 + c = (o[i] + (1L shl (width - 1))) shr width + if (i == 9) o[0] += 19 * c else o[i + 1] += c + o[i] -= c shl width + } + c = (o[0] + (1L shl 25)) shr 26 + o[1] += c + o[0] -= c shl 26 + } + + /** Conditional swap: if b=1, swap p and q element-wise. */ + fun sel25519( + p: LongArray, + q: LongArray, + b: Long, + ) { + val c = b.inv() + 1 // 0 -> 0, 1 -> -1 (all ones) + for (i in 0 until LIMBS) { + val t = c and (p[i] xor q[i]) + p[i] = p[i] xor t + q[i] = q[i] xor t + } + } + + /** + * Encode a field element as 32 little-endian bytes, fully reduced. + * + * This is the only place a canonical representative is produced. The + * leading pass computes the carry that WOULD come out of the top limb if + * the value were >= p, and folds 19 times it back into limb 0; that turns + * any representative of the class — including a non-canonical input such + * as p itself — into the unique one below p. The second pass then carries + * without wrapping, so the top carry falls off and the remaining limbs are + * exactly the base-2^25.5 digits of the answer. + */ + fun pack25519(n: LongArray): ByteArray { + val h = n.copyOf() + var q = (19 * h[9] + (1L shl 24)) shr 25 + for (i in 0 until LIMBS) { + q = (h[i] + q) shr (if (i and 1 == 0) 26 else 25) + } + h[0] += 19 * q + var c = 0L + for (i in 0 until LIMBS) { + val width = if (i and 1 == 0) 26 else 25 + h[i] += c + c = h[i] shr width + h[i] -= c shl width + } + val o = ByteArray(32) + for (i in 0 until LIMBS) { + var v = h[i] + var bit = OFFSET[i] + while (v != 0L) { + val idx = bit shr 3 + val shift = bit and 7 + val room = 8 - shift + o[idx] = (o[idx].toLong() or ((v and ((1L shl room) - 1)) shl shift)).toByte() + v = v ushr room + bit += room } } - for (i in 0 until 15) { - t[i] += 38 * t[i + 16] + return o + } + + /** + * Decode 32 little-endian bytes into a field element. + * + * Bit 255 is ignored, per RFC 7748: the top bit of the last byte is not + * part of the value. Limb 9 is 25 bits wide and ends at bit 254, so the + * mask drops it without a separate step. + */ + fun unpack25519(n: ByteArray): LongArray { + val o = LongArray(LIMBS) + for (i in 0 until LIMBS) { + val width = if (i and 1 == 0) 26 else 25 + val bit = OFFSET[i] + val byteStart = bit shr 3 + val shift = bit and 7 + var chunk = 0L + for (k in 0 until 5) { + val idx = byteStart + k + if (idx < 32) chunk = chunk or ((n[idx].toLong() and 0xFF) shl (8 * k)) + } + o[i] = (chunk ushr shift) and ((1L shl width) - 1) } - for (i in 0 until 16) o[i] = t[i] - car25519(o) - car25519(o) + return o } /** Field inversion: o = a^(-1) (mod p) using Fermat's little theorem. */ fun inv25519(a: LongArray): LongArray { - val o = LongArray(16) - inv25519Into(o, a, LongArray(16), LongArray(31)) + val o = LongArray(LIMBS) + inv25519Into(o, a, LongArray(LIMBS)) 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]. + * a^(p-2) by square-and-multiply over the fixed exponent: 254 squarings + * and ~250 multiplications. [c] is a caller-owned scratch element; [o] may + * alias [a]. */ fun inv25519Into( o: LongArray, a: LongArray, c: LongArray, - t: LongArray, ) { a.copyInto(c) for (i in 253 downTo 0) { - sqrInto(c, c, t) - if (i != 2 && i != 4) mulInto(c, c, a, t) + sqrInto(c, c) + if (i != 2 && i != 4) mulInto(c, c, a) } c.copyInto(o) } @@ -402,16 +577,15 @@ internal object Curve25519Field { /** Parity of a field element (lowest bit after reduction). */ fun par25519(a: LongArray): Int { val d = pack25519(a) - return d[0].toInt() and 1 + return (d[0].toInt() and 1) } /** Raise a field element to the power (2^252 - 3), used in sqrt. */ fun pow2523(a: LongArray): LongArray { val c = a.copyOf() - val t = LongArray(31) for (i in 250 downTo 0) { - sqrInto(c, c, t) - if (i != 1) mulInto(c, c, a, t) + sqrInto(c, c) + if (i != 1) mulInto(c, c, a) } 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 028eddef0a..34fe7fa03b 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 @@ -131,10 +131,10 @@ actual object Ed25519 { private fun newPoint(): Array = arrayOf( - LongArray(16), - LongArray(16), - LongArray(16), - LongArray(16), + LongArray(10), + LongArray(10), + LongArray(10), + LongArray(10), ) private fun identityPoint(): Array { @@ -172,19 +172,16 @@ actual object Ed25519 { * 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) + val a = LongArray(10) + val b = LongArray(10) + val c = LongArray(10) + val d = LongArray(10) + val e = LongArray(10) + val f = LongArray(10) + val g = LongArray(10) + val h = LongArray(10) + val t1 = LongArray(10) + val t2 = LongArray(10) } private fun addPointInPlace( @@ -197,13 +194,13 @@ actual object Ed25519 { // 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.mulInto(s.a, s.a, s.t1) 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.mulInto(s.b, s.b, s.t2) + Curve25519Field.mulInto(s.c, p[3], q[3]) + Curve25519Field.mulInto(s.c, s.c, Curve25519Field.D2) + Curve25519Field.mulInto(s.d, p[2], q[2]) Curve25519Field.addInto(s.d, s.d, s.d) Curve25519Field.subInto(s.e, s.b, s.a) @@ -213,10 +210,10 @@ actual object Ed25519 { // 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) + Curve25519Field.mulInto(p[0], s.e, s.f) + Curve25519Field.mulInto(p[1], s.h, s.g) + Curve25519Field.mulInto(p[2], s.g, s.f) + Curve25519Field.mulInto(p[3], s.e, s.h) } private fun negatePoint(p: Array): Array { @@ -287,25 +284,7 @@ actual object Ed25519 { Curve25519Field.GF1.copyInto(p[2]) val y2 = Curve25519Field.sqr(r) - val d = - Curve25519Field.gf( - 0x78A3, - 0x1359, - 0x4DCA, - 0x75EB, - 0xD8AB, - 0x4141, - 0x0A4D, - 0x0070, - 0xE898, - 0x7779, - 0x4079, - 0x8CC7, - 0xFE73, - 0x2B6F, - 0x6CEE, - 0x5203, - ) + val d = Curve25519Field.D val num = Curve25519Field.sub(y2, Curve25519Field.GF1) val den = Curve25519Field.add(Curve25519Field.mul(d, y2), Curve25519Field.GF1) val denInv = Curve25519Field.inv25519(den) 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 3caeab777f..68d0e477f4 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 @@ -87,17 +87,16 @@ actual object X25519 { // 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) + val e = LongArray(10) + val f = LongArray(10) + val g = LongArray(10) + val h = LongArray(10) + val dd = LongArray(10) + val ff = LongArray(10) + val da = LongArray(10) + val cb = LongArray(10) + val cc = LongArray(10) + val tmp = LongArray(10) for (i in 254 downTo 0) { val r = ((z[i shr 3].toLong() shr (i and 7)) and 1) @@ -111,25 +110,25 @@ actual object X25519 { Curve25519Field.addInto(f, b, d) Curve25519Field.subInto(h, b, d) - Curve25519Field.sqrInto(dd, e, t) - Curve25519Field.sqrInto(ff, g, t) - Curve25519Field.mulInto(da, h, e, t) - Curve25519Field.mulInto(cb, f, g, t) + Curve25519Field.sqrInto(dd, e) + Curve25519Field.sqrInto(ff, g) + Curve25519Field.mulInto(da, h, e) + Curve25519Field.mulInto(cb, f, g) // 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) - Curve25519Field.sqrInto(b, e, t) - Curve25519Field.sqrInto(g, g, t) - Curve25519Field.mulInto(d, g, x, t) + Curve25519Field.sqrInto(b, e) + Curve25519Field.sqrInto(g, g) + Curve25519Field.mulInto(d, g, x) - Curve25519Field.mulInto(a, dd, ff, t) + Curve25519Field.mulInto(a, dd, ff) Curve25519Field.subInto(cc, dd, ff) - Curve25519Field.mulInto(tmp, cc, Curve25519Field.A24, t) + Curve25519Field.mulA24Into(tmp, cc) Curve25519Field.addInto(tmp, dd, tmp) - Curve25519Field.mulInto(c, cc, tmp, t) + Curve25519Field.mulInto(c, cc, tmp) Curve25519Field.sel25519(a, b, r) Curve25519Field.sel25519(c, d, r) @@ -137,8 +136,8 @@ actual object X25519 { // 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) + Curve25519Field.inv25519Into(c, c, tmp) + Curve25519Field.mulInto(a, a, c) 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 1475607b2f..209c36a4b7 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 @@ -131,10 +131,10 @@ actual object Ed25519 { private fun newPoint(): Array = arrayOf( - LongArray(16), - LongArray(16), - LongArray(16), - LongArray(16), + LongArray(10), + LongArray(10), + LongArray(10), + LongArray(10), ) private fun identityPoint(): Array { @@ -172,19 +172,16 @@ actual object Ed25519 { * 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) + val a = LongArray(10) + val b = LongArray(10) + val c = LongArray(10) + val d = LongArray(10) + val e = LongArray(10) + val f = LongArray(10) + val g = LongArray(10) + val h = LongArray(10) + val t1 = LongArray(10) + val t2 = LongArray(10) } private fun addPointInPlace( @@ -197,13 +194,13 @@ actual object Ed25519 { // 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.mulInto(s.a, s.a, s.t1) 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.mulInto(s.b, s.b, s.t2) + Curve25519Field.mulInto(s.c, p[3], q[3]) + Curve25519Field.mulInto(s.c, s.c, Curve25519Field.D2) + Curve25519Field.mulInto(s.d, p[2], q[2]) Curve25519Field.addInto(s.d, s.d, s.d) Curve25519Field.subInto(s.e, s.b, s.a) @@ -213,10 +210,10 @@ actual object Ed25519 { // 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) + Curve25519Field.mulInto(p[0], s.e, s.f) + Curve25519Field.mulInto(p[1], s.h, s.g) + Curve25519Field.mulInto(p[2], s.g, s.f) + Curve25519Field.mulInto(p[3], s.e, s.h) } private fun negatePoint(p: Array): Array { @@ -287,25 +284,7 @@ actual object Ed25519 { Curve25519Field.GF1.copyInto(p[2]) val y2 = Curve25519Field.sqr(r) - val d = - Curve25519Field.gf( - 0x78A3, - 0x1359, - 0x4DCA, - 0x75EB, - 0xD8AB, - 0x4141, - 0x0A4D, - 0x0070, - 0xE898, - 0x7779, - 0x4079, - 0x8CC7, - 0xFE73, - 0x2B6F, - 0x6CEE, - 0x5203, - ) + val d = Curve25519Field.D val num = Curve25519Field.sub(y2, Curve25519Field.GF1) val den = Curve25519Field.add(Curve25519Field.mul(d, y2), Curve25519Field.GF1) val denInv = Curve25519Field.inv25519(den) 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 ed0d6a7ffa..413eb64faa 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 @@ -82,17 +82,16 @@ actual object X25519 { // 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) + val e = LongArray(10) + val f = LongArray(10) + val g = LongArray(10) + val h = LongArray(10) + val dd = LongArray(10) + val ff = LongArray(10) + val da = LongArray(10) + val cb = LongArray(10) + val cc = LongArray(10) + val tmp = LongArray(10) for (i in 254 downTo 0) { val r = ((z[i shr 3].toLong() shr (i and 7)) and 1) @@ -106,25 +105,25 @@ actual object X25519 { Curve25519Field.addInto(f, b, d) Curve25519Field.subInto(h, b, d) - Curve25519Field.sqrInto(dd, e, t) - Curve25519Field.sqrInto(ff, g, t) - Curve25519Field.mulInto(da, h, e, t) - Curve25519Field.mulInto(cb, f, g, t) + Curve25519Field.sqrInto(dd, e) + Curve25519Field.sqrInto(ff, g) + Curve25519Field.mulInto(da, h, e) + Curve25519Field.mulInto(cb, f, g) // 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) - Curve25519Field.sqrInto(b, e, t) - Curve25519Field.sqrInto(g, g, t) - Curve25519Field.mulInto(d, g, x, t) + Curve25519Field.sqrInto(b, e) + Curve25519Field.sqrInto(g, g) + Curve25519Field.mulInto(d, g, x) - Curve25519Field.mulInto(a, dd, ff, t) + Curve25519Field.mulInto(a, dd, ff) Curve25519Field.subInto(cc, dd, ff) - Curve25519Field.mulInto(tmp, cc, Curve25519Field.A24, t) + Curve25519Field.mulA24Into(tmp, cc) Curve25519Field.addInto(tmp, dd, tmp) - Curve25519Field.mulInto(c, cc, tmp, t) + Curve25519Field.mulInto(c, cc, tmp) Curve25519Field.sel25519(a, b, r) Curve25519Field.sel25519(c, d, r) @@ -132,8 +131,8 @@ actual object X25519 { // 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) + Curve25519Field.inv25519Into(c, c, tmp) + Curve25519Field.mulInto(a, a, c) return Curve25519Field.pack25519(a) } }