From 3977c38fedf96b1d520f8acf6a21aa15e2d1c719 Mon Sep 17 00:00:00 2001 From: Claude Date: Thu, 10 Sep 2026 14:32:55 +0000 Subject: [PATCH] perf(marmot): field arithmetic on 10 limbs of radix 2^25.5 MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit One X25519 scalar multiplication cost 541us, and create_group/1 is about two dozen of them, so the curve primitive was not part of the gap against MDK — it was the gap. The cause was the representation rather than the language. Curve25519Field used TweetNaCl's 16 limbs of radix 2^16, so a schoolbook multiply spent 256 limb products. SunEC's X25519 — also pure Java, same JIT, same machine — ran the same operation in 160us on ~26-bit limbs in 10 words, which is 100 products. The ratio of products matched the ratio of times, which rules out "managed language" as the explanation and names the fix. Rewrite the field to 10 limbs of radix 2^25.5, the layout ref10, curve25519-donna and SunEC all use. A multiply is 100 products, a square 55 (each off-diagonal pair once, doubled), and the ladder's a24 constant gets a dedicated scalar multiply instead of a general one against nine zero limbs. Limbs stay signed and denormalised between operations; only pack25519 produces a canonical value. Straight-line locals mean mulInto and sqrInto need no scratch accumulator at all, so that parameter is gone from every caller. x25519_dh 541us -> 121us ed25519_sign 1018us -> 259us x25519_base 535us -> 121us ed25519_verify 2113us -> 536us At 121us the scalar multiplication is faster than SunEC's 160us, which is the sanity check on the result: it lands where a good managed implementation should rather than somewhere suspiciously better. Against MDK, create_group/1 goes from 4.7x slower to 1.5x, join_welcome from 1.3x slower to 2.5x FASTER, and send_app_message from 2.5x to 6.5x faster. Re-encoding curve constants is where one mistyped limb yields code that runs and is silently wrong, so none were transcribed by hand. Each was re-derived from its existing encoding and checked against its mathematical definition: d == -121665/121666, d2 == 2d, By == 4/5, I^2 == -1. The multiply and square formulas were generated from the representation's weight bookkeeping and diffed against an independent reference over 20000 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. RFC 7748, RFC 8032, HPKE, the MDK crypto-interop vectors and the full quartz and commons suites all pass unchanged. Co-Authored-By: Claude Opus 5 Claude-Session: https://claude.ai/code/session_016kCuA6tc4JQzHPCDd39GHq --- marmotBench/README.md | 61 ++ .../quartz/marmot/mls/crypto/Ed25519.apple.kt | 71 +- .../quartz/marmot/mls/crypto/X25519.apple.kt | 45 +- .../marmot/mls/crypto/Curve25519Field.kt | 740 +++++++++++------- .../marmot/mls/crypto/Ed25519.jvmAndroid.kt | 69 +- .../marmot/mls/crypto/X25519.jvmAndroid.kt | 45 +- .../quartz/marmot/mls/crypto/Ed25519.linux.kt | 69 +- .../quartz/marmot/mls/crypto/X25519.linux.kt | 45 +- 8 files changed, 657 insertions(+), 488 deletions(-) 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) } }