Merge pull request #2178 from vitorpamplona/claude/complete-marmot-ios-dM5C0

Implement Ed25519 and X25519 cryptography for native platforms
This commit is contained in:
Vitor Pamplona
2026-04-08 10:32:06 -04:00
committed by GitHub
8 changed files with 1240 additions and 28 deletions
@@ -20,19 +20,421 @@
*/
package com.vitorpamplona.quartz.marmot.mls.crypto
import com.vitorpamplona.quartz.utils.RandomInstance
import io.github.andreypfau.kotlinx.crypto.Sha512
/**
* Apple/Native Ed25519 implementation using pure Kotlin field arithmetic.
*
* Implements RFC 8032 Ed25519 digital signatures.
* Private key format: 32-byte seed + 32-byte public key (64 bytes total).
* Public key format: 32-byte compressed Edwards point.
*
* Based on the TweetNaCl algorithm by Bernstein et al.
*/
actual object Ed25519 {
actual fun generateKeyPair(): Ed25519KeyPair = TODO("Ed25519 not yet implemented for this platform")
private const val SEED_LENGTH = 32
private const val PUBLIC_KEY_LENGTH = 32
actual fun generateKeyPair(): Ed25519KeyPair {
val seed = RandomInstance.bytes(SEED_LENGTH)
val publicKey = derivePublicKey(seed)
val privateKey = seed + publicKey
return Ed25519KeyPair(privateKey, publicKey)
}
actual fun sign(
message: ByteArray,
privateKey: ByteArray,
): ByteArray = TODO("Ed25519 not yet implemented for this platform")
): ByteArray {
require(privateKey.size == SEED_LENGTH * 2) { "Private key must be 64 bytes (seed + public)" }
val seed = privateKey.copyOfRange(0, SEED_LENGTH)
val publicKey = privateKey.copyOfRange(SEED_LENGTH, SEED_LENGTH * 2)
// SHA-512(seed) -> expanded key
val d = sha512(seed)
// Clamp the scalar
d[0] = (d[0].toInt() and 248).toByte()
d[31] = ((d[31].toInt() and 63) or 64).toByte()
// r = SHA-512(d[32..63] || message) mod L
val rHash = sha512(d.copyOfRange(32, 64) + message)
val r = reduce(rHash)
// R = [r]B
val rPoint = scalarMultBase(r)
val rBytes = packPoint(rPoint)
// S = (r + SHA-512(R || publicKey || message) * s) mod L
val hramHash = sha512(rBytes + publicKey + message)
val hram = reduce(hramHash)
val signature = ByteArray(64)
rBytes.copyInto(signature, 0)
// Compute s = r + hram * d (mod L) in 64-byte arithmetic
val x = LongArray(64)
for (i in 0 until 32) x[i] = r[i].toLong() and 0xFF
for (i in 0 until 32) {
for (j in 0 until 32) {
x[i + j] += (hram[i].toLong() and 0xFF) * (d[j].toLong() and 0xFF)
}
}
val sBytes = modL(x)
sBytes.copyInto(signature, 32)
return signature
}
actual fun verify(
message: ByteArray,
signature: ByteArray,
publicKey: ByteArray,
): Boolean = TODO("Ed25519 not yet implemented for this platform")
): Boolean {
require(publicKey.size == PUBLIC_KEY_LENGTH) { "Public key must be 32 bytes" }
if (signature.size != 64) return false
actual fun publicFromPrivate(privateKey: ByteArray): ByteArray = TODO("Ed25519 not yet implemented for this platform")
val aPoint = unpackPoint(publicKey) ?: return false
val rBytes = signature.copyOfRange(0, 32)
val sBytes = signature.copyOfRange(32, 64)
// Check s < L
if (!isCanonicalScalar(sBytes)) return false
val hramHash = sha512(rBytes + publicKey + message)
val hram = reduce(hramHash)
// Verify: [s]B = R + [hram]A
// Equivalent: [s]B - [hram]A = R
// We compute: [s]B and [hram]A separately, then check
val sPoint = scalarMultBase(sBytes)
// Negate A for subtraction: compute [-hram]A
val hramA = scalarMult(aPoint, hram)
// R_check = [s]B - [hram]A = [s]B + [-hram]A
val negHramA = negatePoint(hramA)
val rCheck = addPoints(sPoint, negHramA)
val rCheckBytes = packPoint(rCheck)
return rCheckBytes.contentEquals(rBytes)
}
actual fun publicFromPrivate(privateKey: ByteArray): ByteArray {
require(privateKey.size == SEED_LENGTH * 2) { "Private key must be 64 bytes (seed + public)" }
return privateKey.copyOfRange(SEED_LENGTH, SEED_LENGTH * 2)
}
// --- Internal operations ---
/** Derive Ed25519 public key from 32-byte seed. */
private fun derivePublicKey(seed: ByteArray): ByteArray {
val d = sha512(seed)
d[0] = (d[0].toInt() and 248).toByte()
d[31] = ((d[31].toInt() and 63) or 64).toByte()
val p = scalarMultBase(d.copyOfRange(0, 32))
return packPoint(p)
}
/** SHA-512 hash using the platform library. */
private fun sha512(data: ByteArray): ByteArray {
val digest = Sha512()
digest.update(data)
return digest.digest()
}
// --- Extended Edwards point operations ---
// Point = Array of 4 LongArray(16), representing (X, Y, Z, T)
// where x = X/Z, y = Y/Z, x*y = T/Z
private fun newPoint(): Array<LongArray> =
arrayOf(
LongArray(16),
LongArray(16),
LongArray(16),
LongArray(16),
)
/** Set point to the identity (0, 1, 1, 0). */
private fun identityPoint(): Array<LongArray> {
val p = newPoint()
p[1][0] = 1
p[2][0] = 1
return p
}
/** Point addition on extended twisted Edwards curve. */
private fun addPoints(
p: Array<LongArray>,
q: Array<LongArray>,
): Array<LongArray> {
val result =
arrayOf(
p[0].copyOf(),
p[1].copyOf(),
p[2].copyOf(),
p[3].copyOf(),
)
addPointInPlace(result, q)
return result
}
/** In-place point addition: p += q. */
private fun addPointInPlace(
p: Array<LongArray>,
q: Array<LongArray>,
) {
val a = Curve25519Field.sub(p[1], p[0])
val t = Curve25519Field.sub(q[1], q[0])
val aMul = Curve25519Field.mul(a, t)
val b = Curve25519Field.add(p[0], p[1])
val t2 = Curve25519Field.add(q[0], q[1])
val bMul = Curve25519Field.mul(b, t2)
val c = Curve25519Field.mul(p[3], q[3])
val cMul = Curve25519Field.mul(c, Curve25519Field.D2)
val d = Curve25519Field.mul(p[2], q[2])
val dAdd = Curve25519Field.add(d, d)
val e = Curve25519Field.sub(bMul, aMul)
val f = Curve25519Field.sub(dAdd, cMul)
val g = Curve25519Field.add(dAdd, cMul)
val h = Curve25519Field.add(bMul, aMul)
Curve25519Field.mul(e, f).copyInto(p[0])
Curve25519Field.mul(h, g).copyInto(p[1])
Curve25519Field.mul(g, f).copyInto(p[2])
Curve25519Field.mul(e, h).copyInto(p[3])
}
/** Point doubling (self-addition). */
private fun doublePoint(p: Array<LongArray>): Array<LongArray> = addPoints(p, p)
/** Negate a point: (X, Y, Z, T) -> (-X, Y, Z, -T). */
private fun negatePoint(p: Array<LongArray>): Array<LongArray> {
val result = newPoint()
Curve25519Field.sub(Curve25519Field.GF0, p[0]).copyInto(result[0])
p[1].copyInto(result[1])
p[2].copyInto(result[2])
Curve25519Field.sub(Curve25519Field.GF0, p[3]).copyInto(result[3])
return result
}
/** Scalar multiplication: [s]P using double-and-add. */
private fun scalarMult(
p: Array<LongArray>,
s: ByteArray,
): Array<LongArray> {
val result = identityPoint()
val q =
arrayOf(
p[0].copyOf(),
p[1].copyOf(),
p[2].copyOf(),
p[3].copyOf(),
)
for (i in 255 downTo 0) {
val b = ((s[i shr 3].toInt() shr (i and 7)) and 1).toLong()
cswap(result, q, b)
addPointInPlace(q, result)
addPointInPlace(result, result)
cswap(result, q, b)
}
return result
}
/** Scalar multiplication with the base point: [s]B. */
private fun scalarMultBase(s: ByteArray): Array<LongArray> {
val basePoint = newPoint()
Curve25519Field.BX.copyInto(basePoint[0])
Curve25519Field.BY.copyInto(basePoint[1])
Curve25519Field.GF1.copyInto(basePoint[2])
Curve25519Field.mul(Curve25519Field.BX, Curve25519Field.BY).copyInto(basePoint[3])
return scalarMult(basePoint, s)
}
/** Conditional swap of two points. */
private fun cswap(
p: Array<LongArray>,
q: Array<LongArray>,
b: Long,
) {
for (i in 0 until 4) {
Curve25519Field.sel25519(p[i], q[i], b)
}
}
/** Pack an extended Edwards point to 32-byte compressed encoding. */
private fun packPoint(p: Array<LongArray>): ByteArray {
val zi = Curve25519Field.inv25519(p[2])
val tx = Curve25519Field.mul(p[0], zi)
val ty = Curve25519Field.mul(p[1], zi)
val r = Curve25519Field.pack25519(ty)
r[31] = (r[31].toInt() xor (Curve25519Field.par25519(tx) shl 7)).toByte()
return r
}
/**
* Unpack a 32-byte compressed Edwards point.
* Returns null if the point is not on the curve.
*/
private fun unpackPoint(s: ByteArray): Array<LongArray>? {
val p = newPoint()
val r = Curve25519Field.unpack25519(s)
r.copyInto(p[1])
Curve25519Field.GF1.copyInto(p[2])
// 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 num = Curve25519Field.sub(y2, Curve25519Field.GF1)
val den = Curve25519Field.add(Curve25519Field.mul(d, y2), Curve25519Field.GF1)
val denInv = Curve25519Field.inv25519(den)
var x2 = Curve25519Field.mul(num, denInv)
// Try sqrt(x2)
var x = Curve25519Field.pow2523(x2)
x = Curve25519Field.mul(x, x2)
// Check: x^2 == x2?
val check = Curve25519Field.sub(Curve25519Field.sqr(x), x2)
val checkPacked = Curve25519Field.pack25519(check)
if (!checkPacked.all { it == 0.toByte() }) {
// Try x * sqrt(-1)
x = Curve25519Field.mul(x, Curve25519Field.I)
val check2 = Curve25519Field.sub(Curve25519Field.sqr(x), x2)
val check2Packed = Curve25519Field.pack25519(check2)
if (!check2Packed.all { it == 0.toByte() }) {
return null
}
}
// Adjust sign
if (Curve25519Field.par25519(x) != ((s[31].toInt() shr 7) and 1)) {
x = Curve25519Field.sub(Curve25519Field.GF0, x)
}
x.copyInto(p[0])
Curve25519Field.mul(p[0], p[1]).copyInto(p[3])
return p
}
// --- Scalar reduction modulo L ---
// L = 2^252 + 27742317777372353535851937790883648493
private val L =
longArrayOf(
0xED,
0xD3,
0xF5,
0x5C,
0x1A,
0x63,
0x12,
0x58,
0xD6,
0x9C,
0xF7,
0xA2,
0xDE,
0xF9,
0xDE,
0x14,
0,
0,
0,
0,
0,
0,
0,
0,
0,
0,
0,
0,
0,
0,
0,
0x10,
)
/**
* Reduce a 64-byte hash to a 32-byte scalar modulo L.
* This is used for Ed25519 nonce and challenge computation.
*/
private fun reduce(input: ByteArray): ByteArray {
val x = LongArray(64)
for (i in 0 until 64) x[i] = input[i].toLong() and 0xFF
return modL(x)
}
/**
* Reduce a 64-element long array modulo L, returning a 32-byte result.
*/
private fun modL(x: LongArray): ByteArray {
for (i in 63 downTo 32) {
var carry: Long = 0
var j = i - 32
val k = i - 12
while (j < k) {
x[j] += carry - 16 * x[i] * L[j - (i - 32)]
carry = (x[j] + 128) shr 8
x[j] -= carry shl 8
j++
}
x[j] += carry
x[i] = 0
}
var carry: Long = 0
for (j in 0 until 32) {
x[j] += carry - (x[31] shr 4) * L[j]
carry = x[j] shr 8
x[j] = x[j] and 0xFF
}
for (j in 0 until 32) {
x[j] -= carry * L[j]
}
val r = ByteArray(32)
for (i in 0 until 32) {
x[i + 1] += x[i] shr 8
r[i] = (x[i] and 0xFF).toByte()
}
return r
}
/** Check if a scalar s < L (canonical). */
private fun isCanonicalScalar(s: ByteArray): Boolean {
// Check s < L by comparing from high byte to low
var borrow: Long = 0
for (i in 31 downTo 0) {
val si = s[i].toLong() and 0xFF
val li = L[i]
if (si < li + borrow) return true
if (si > li + borrow) return false
borrow = 0
}
return false // s == L is not canonical
}
}
@@ -20,13 +20,106 @@
*/
package com.vitorpamplona.quartz.marmot.mls.crypto
import com.vitorpamplona.quartz.utils.RandomInstance
/**
* Apple/Native X25519 implementation using pure Kotlin field arithmetic.
*
* Implements RFC 7748 X25519 Diffie-Hellman key agreement via Montgomery ladder.
* Key format: raw 32-byte Curve25519 keys (little-endian per RFC 7748).
*/
actual object X25519 {
actual fun generateKeyPair(): X25519KeyPair = TODO("X25519 not yet implemented for this platform")
private const val KEY_LENGTH = 32
actual fun generateKeyPair(): X25519KeyPair {
val privateKey = RandomInstance.bytes(KEY_LENGTH)
val publicKey = publicFromPrivate(privateKey)
return X25519KeyPair(privateKey, publicKey)
}
actual fun dh(
privateKey: ByteArray,
publicKey: ByteArray,
): ByteArray = TODO("X25519 not yet implemented for this platform")
): ByteArray {
require(privateKey.size == KEY_LENGTH) { "Private key must be 32 bytes" }
require(publicKey.size == KEY_LENGTH) { "Public key must be 32 bytes" }
actual fun publicFromPrivate(privateKey: ByteArray): ByteArray = TODO("X25519 not yet implemented for this platform")
val result = scalarmult(privateKey, publicKey)
require(!result.all { it == 0.toByte() }) {
"DH produced all-zero shared secret (possible small-subgroup attack)"
}
return result
}
actual fun publicFromPrivate(privateKey: ByteArray): ByteArray {
require(privateKey.size == KEY_LENGTH) { "Private key must be 32 bytes" }
val basepoint = ByteArray(KEY_LENGTH)
basepoint[0] = 9
return scalarmult(privateKey, basepoint)
}
/**
* X25519 scalar multiplication via Montgomery ladder (RFC 7748).
*
* Computes [n]P on Curve25519 in Montgomery form.
* Based on the TweetNaCl algorithm by Bernstein et al.
*/
private fun scalarmult(
n: ByteArray,
p: ByteArray,
): ByteArray {
val z = n.copyOf()
// Clamp scalar per RFC 7748 Section 5
z[0] = (z[0].toInt() and 248).toByte()
z[31] = ((z[31].toInt() and 127) or 64).toByte()
val x = Curve25519Field.unpack25519(p)
val a = Curve25519Field.GF1.copyOf()
val b = x.copyOf()
val c = Curve25519Field.GF0.copyOf()
val d = Curve25519Field.GF1.copyOf()
for (i in 254 downTo 0) {
val r = ((z[i shr 3].toLong() shr (i and 7)) and 1)
Curve25519Field.sel25519(a, b, r)
Curve25519Field.sel25519(c, d, r)
val e = Curve25519Field.add(a, c)
val aMc = Curve25519Field.sub(a, c)
val f = Curve25519Field.add(b, d)
val bMd = Curve25519Field.sub(b, d)
val dd = Curve25519Field.sqr(e)
val ff = Curve25519Field.sqr(aMc)
val da = Curve25519Field.mul(bMd, e)
val cb = Curve25519Field.mul(f, aMc)
val ePrime = Curve25519Field.add(da, cb)
val aPrime = Curve25519Field.sub(da, cb)
val bNew = Curve25519Field.sqr(ePrime)
val aSqr = Curve25519Field.sqr(aPrime)
val dNew = Curve25519Field.mul(aSqr, x)
val aNew = Curve25519Field.mul(dd, ff)
val cc = Curve25519Field.sub(dd, ff)
val tmp = Curve25519Field.mul(cc, Curve25519Field.A24)
val ddPlusTmp = Curve25519Field.add(dd, tmp)
val cNew = Curve25519Field.mul(cc, ddPlusTmp)
aNew.copyInto(a)
bNew.copyInto(b)
cNew.copyInto(c)
dNew.copyInto(d)
Curve25519Field.sel25519(a, b, r)
Curve25519Field.sel25519(c, d, r)
}
val invC = Curve25519Field.inv25519(c)
val result = Curve25519Field.mul(a, invC)
return Curve25519Field.pack25519(result)
}
}
@@ -637,15 +637,28 @@ class MlsGroup private constructor(
senderLeafIndex: Int,
confirmationTag: ByteArray? = null,
) {
require(senderLeafIndex >= 0 && senderLeafIndex < tree.leafCount) {
"Invalid sender leaf index: $senderLeafIndex"
}
require(tree.getLeaf(senderLeafIndex) != null) {
"Sender leaf is blank at index $senderLeafIndex"
}
val commit = Commit.decodeTls(TlsReader(commitBytes))
// External commits (containing ExternalInit) have a sender that is not
// yet in the tree — their leaf will be added via the UpdatePath below.
val isExternalCommit =
commit.proposals.any {
it is ProposalOrRef.Inline && it.proposal is Proposal.ExternalInit
}
if (isExternalCommit) {
require(senderLeafIndex >= 0 && senderLeafIndex <= tree.leafCount) {
"Invalid sender leaf index for external commit: $senderLeafIndex"
}
} else {
require(senderLeafIndex >= 0 && senderLeafIndex < tree.leafCount) {
"Invalid sender leaf index: $senderLeafIndex"
}
require(tree.getLeaf(senderLeafIndex) != null) {
"Sender leaf is blank at index $senderLeafIndex"
}
}
// Apply proposals (resolve references from pending pool)
for (proposalOrRef in commit.proposals) {
when (proposalOrRef) {
@@ -122,7 +122,7 @@ data class MlsKeyPackage(
val version = reader.readUint16()
require(version == 1) { "Unsupported MLS version: $version" }
val cipherSuite = reader.readUint16()
require(cipherSuite == 1) { "Unsupported ciphersuite: $cipherSuite" }
require(cipherSuite in 1..0xFFFF) { "Invalid ciphersuite: $cipherSuite" }
return MlsKeyPackage(
version = version,
cipherSuite = cipherSuite,
@@ -142,24 +142,24 @@ class MarmotSubscriptionManagerTest {
}
@Test
fun testBuildFiltersIncludesBothTypes() =
fun testBuildFiltersIncludesAllTypes() =
runTest {
val manager = MarmotSubscriptionManager(userPubKey)
manager.subscribeGroup(groupId1)
val allFilters = manager.buildFilters()
// Should have 1 group filter + 1 gift wrap filter
assertEquals(2, allFilters.size)
// Should have 1 group filter + 1 gift wrap filter + 1 own key package filter
assertEquals(3, allFilters.size)
}
@Test
fun testBuildFiltersWithNoGroupsHasGiftWrapOnly() {
fun testBuildFiltersWithNoGroupsHasGiftWrapAndKeyPackage() {
val manager = MarmotSubscriptionManager(userPubKey)
val allFilters = manager.buildFilters()
// Only the gift wrap filter
assertEquals(1, allFilters.size)
// Gift wrap filter + own key package filter
assertEquals(2, allFilters.size)
assertEquals(listOf(GiftWrapEvent.KIND), allFilters[0].kinds)
}
@@ -20,19 +20,367 @@
*/
package com.vitorpamplona.quartz.marmot.mls.crypto
import com.vitorpamplona.quartz.utils.RandomInstance
import io.github.andreypfau.kotlinx.crypto.Sha512
/**
* Linux/Native Ed25519 implementation using pure Kotlin field arithmetic.
*
* Implements RFC 8032 Ed25519 digital signatures.
* Private key format: 32-byte seed + 32-byte public key (64 bytes total).
* Public key format: 32-byte compressed Edwards point.
*
* Based on the TweetNaCl algorithm by Bernstein et al.
*/
actual object Ed25519 {
actual fun generateKeyPair(): Ed25519KeyPair = TODO("Ed25519 not yet implemented for this platform")
private const val SEED_LENGTH = 32
private const val PUBLIC_KEY_LENGTH = 32
actual fun generateKeyPair(): Ed25519KeyPair {
val seed = RandomInstance.bytes(SEED_LENGTH)
val publicKey = derivePublicKey(seed)
val privateKey = seed + publicKey
return Ed25519KeyPair(privateKey, publicKey)
}
actual fun sign(
message: ByteArray,
privateKey: ByteArray,
): ByteArray = TODO("Ed25519 not yet implemented for this platform")
): ByteArray {
require(privateKey.size == SEED_LENGTH * 2) { "Private key must be 64 bytes (seed + public)" }
val seed = privateKey.copyOfRange(0, SEED_LENGTH)
val publicKey = privateKey.copyOfRange(SEED_LENGTH, SEED_LENGTH * 2)
val d = sha512(seed)
d[0] = (d[0].toInt() and 248).toByte()
d[31] = ((d[31].toInt() and 63) or 64).toByte()
val rHash = sha512(d.copyOfRange(32, 64) + message)
val r = reduce(rHash)
val rPoint = scalarMultBase(r)
val rBytes = packPoint(rPoint)
val hramHash = sha512(rBytes + publicKey + message)
val hram = reduce(hramHash)
val signature = ByteArray(64)
rBytes.copyInto(signature, 0)
val x = LongArray(64)
for (i in 0 until 32) x[i] = r[i].toLong() and 0xFF
for (i in 0 until 32) {
for (j in 0 until 32) {
x[i + j] += (hram[i].toLong() and 0xFF) * (d[j].toLong() and 0xFF)
}
}
val sBytes = modL(x)
sBytes.copyInto(signature, 32)
return signature
}
actual fun verify(
message: ByteArray,
signature: ByteArray,
publicKey: ByteArray,
): Boolean = TODO("Ed25519 not yet implemented for this platform")
): Boolean {
require(publicKey.size == PUBLIC_KEY_LENGTH) { "Public key must be 32 bytes" }
if (signature.size != 64) return false
actual fun publicFromPrivate(privateKey: ByteArray): ByteArray = TODO("Ed25519 not yet implemented for this platform")
val aPoint = unpackPoint(publicKey) ?: return false
val rBytes = signature.copyOfRange(0, 32)
val sBytes = signature.copyOfRange(32, 64)
if (!isCanonicalScalar(sBytes)) return false
val hramHash = sha512(rBytes + publicKey + message)
val hram = reduce(hramHash)
val sPoint = scalarMultBase(sBytes)
val hramA = scalarMult(aPoint, hram)
val negHramA = negatePoint(hramA)
val rCheck = addPoints(sPoint, negHramA)
val rCheckBytes = packPoint(rCheck)
return rCheckBytes.contentEquals(rBytes)
}
actual fun publicFromPrivate(privateKey: ByteArray): ByteArray {
require(privateKey.size == SEED_LENGTH * 2) { "Private key must be 64 bytes (seed + public)" }
return privateKey.copyOfRange(SEED_LENGTH, SEED_LENGTH * 2)
}
private fun derivePublicKey(seed: ByteArray): ByteArray {
val d = sha512(seed)
d[0] = (d[0].toInt() and 248).toByte()
d[31] = ((d[31].toInt() and 63) or 64).toByte()
val p = scalarMultBase(d.copyOfRange(0, 32))
return packPoint(p)
}
private fun sha512(data: ByteArray): ByteArray {
val digest = Sha512()
digest.update(data)
return digest.digest()
}
private fun newPoint(): Array<LongArray> =
arrayOf(
LongArray(16),
LongArray(16),
LongArray(16),
LongArray(16),
)
private fun identityPoint(): Array<LongArray> {
val p = newPoint()
p[1][0] = 1
p[2][0] = 1
return p
}
private fun addPoints(
p: Array<LongArray>,
q: Array<LongArray>,
): Array<LongArray> {
val result =
arrayOf(
p[0].copyOf(),
p[1].copyOf(),
p[2].copyOf(),
p[3].copyOf(),
)
addPointInPlace(result, q)
return result
}
private fun addPointInPlace(
p: Array<LongArray>,
q: Array<LongArray>,
) {
val a = Curve25519Field.sub(p[1], p[0])
val t = Curve25519Field.sub(q[1], q[0])
val aMul = Curve25519Field.mul(a, t)
val b = Curve25519Field.add(p[0], p[1])
val t2 = Curve25519Field.add(q[0], q[1])
val bMul = Curve25519Field.mul(b, t2)
val c = Curve25519Field.mul(p[3], q[3])
val cMul = Curve25519Field.mul(c, Curve25519Field.D2)
val d = Curve25519Field.mul(p[2], q[2])
val dAdd = Curve25519Field.add(d, d)
val e = Curve25519Field.sub(bMul, aMul)
val f = Curve25519Field.sub(dAdd, cMul)
val g = Curve25519Field.add(dAdd, cMul)
val h = Curve25519Field.add(bMul, aMul)
Curve25519Field.mul(e, f).copyInto(p[0])
Curve25519Field.mul(h, g).copyInto(p[1])
Curve25519Field.mul(g, f).copyInto(p[2])
Curve25519Field.mul(e, h).copyInto(p[3])
}
private fun negatePoint(p: Array<LongArray>): Array<LongArray> {
val result = newPoint()
Curve25519Field.sub(Curve25519Field.GF0, p[0]).copyInto(result[0])
p[1].copyInto(result[1])
p[2].copyInto(result[2])
Curve25519Field.sub(Curve25519Field.GF0, p[3]).copyInto(result[3])
return result
}
private fun scalarMult(
p: Array<LongArray>,
s: ByteArray,
): Array<LongArray> {
val result = identityPoint()
val q =
arrayOf(
p[0].copyOf(),
p[1].copyOf(),
p[2].copyOf(),
p[3].copyOf(),
)
for (i in 255 downTo 0) {
val b = ((s[i shr 3].toInt() shr (i and 7)) and 1).toLong()
cswap(result, q, b)
addPointInPlace(q, result)
addPointInPlace(result, result)
cswap(result, q, b)
}
return result
}
private fun scalarMultBase(s: ByteArray): Array<LongArray> {
val basePoint = newPoint()
Curve25519Field.BX.copyInto(basePoint[0])
Curve25519Field.BY.copyInto(basePoint[1])
Curve25519Field.GF1.copyInto(basePoint[2])
Curve25519Field.mul(Curve25519Field.BX, Curve25519Field.BY).copyInto(basePoint[3])
return scalarMult(basePoint, s)
}
private fun cswap(
p: Array<LongArray>,
q: Array<LongArray>,
b: Long,
) {
for (i in 0 until 4) {
Curve25519Field.sel25519(p[i], q[i], b)
}
}
private fun packPoint(p: Array<LongArray>): ByteArray {
val zi = Curve25519Field.inv25519(p[2])
val tx = Curve25519Field.mul(p[0], zi)
val ty = Curve25519Field.mul(p[1], zi)
val r = Curve25519Field.pack25519(ty)
r[31] = (r[31].toInt() xor (Curve25519Field.par25519(tx) shl 7)).toByte()
return r
}
private fun unpackPoint(s: ByteArray): Array<LongArray>? {
val p = newPoint()
val r = Curve25519Field.unpack25519(s)
r.copyInto(p[1])
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 num = Curve25519Field.sub(y2, Curve25519Field.GF1)
val den = Curve25519Field.add(Curve25519Field.mul(d, y2), Curve25519Field.GF1)
val denInv = Curve25519Field.inv25519(den)
var x2 = Curve25519Field.mul(num, denInv)
var x = Curve25519Field.pow2523(x2)
x = Curve25519Field.mul(x, x2)
val check = Curve25519Field.sub(Curve25519Field.sqr(x), x2)
val checkPacked = Curve25519Field.pack25519(check)
if (!checkPacked.all { it == 0.toByte() }) {
x = Curve25519Field.mul(x, Curve25519Field.I)
val check2 = Curve25519Field.sub(Curve25519Field.sqr(x), x2)
val check2Packed = Curve25519Field.pack25519(check2)
if (!check2Packed.all { it == 0.toByte() }) {
return null
}
}
if (Curve25519Field.par25519(x) != ((s[31].toInt() shr 7) and 1)) {
x = Curve25519Field.sub(Curve25519Field.GF0, x)
}
x.copyInto(p[0])
Curve25519Field.mul(p[0], p[1]).copyInto(p[3])
return p
}
private val L =
longArrayOf(
0xED,
0xD3,
0xF5,
0x5C,
0x1A,
0x63,
0x12,
0x58,
0xD6,
0x9C,
0xF7,
0xA2,
0xDE,
0xF9,
0xDE,
0x14,
0,
0,
0,
0,
0,
0,
0,
0,
0,
0,
0,
0,
0,
0,
0,
0x10,
)
private fun reduce(input: ByteArray): ByteArray {
val x = LongArray(64)
for (i in 0 until 64) x[i] = input[i].toLong() and 0xFF
return modL(x)
}
private fun modL(x: LongArray): ByteArray {
for (i in 63 downTo 32) {
var carry: Long = 0
var j = i - 32
val k = i - 12
while (j < k) {
x[j] += carry - 16 * x[i] * L[j - (i - 32)]
carry = (x[j] + 128) shr 8
x[j] -= carry shl 8
j++
}
x[j] += carry
x[i] = 0
}
var carry: Long = 0
for (j in 0 until 32) {
x[j] += carry - (x[31] shr 4) * L[j]
carry = x[j] shr 8
x[j] = x[j] and 0xFF
}
for (j in 0 until 32) {
x[j] -= carry * L[j]
}
val r = ByteArray(32)
for (i in 0 until 32) {
x[i + 1] += x[i] shr 8
r[i] = (x[i] and 0xFF).toByte()
}
return r
}
private fun isCanonicalScalar(s: ByteArray): Boolean {
var borrow: Long = 0
for (i in 31 downTo 0) {
val si = s[i].toLong() and 0xFF
val li = L[i]
if (si < li + borrow) return true
if (si > li + borrow) return false
borrow = 0
}
return false
}
}
@@ -20,13 +20,102 @@
*/
package com.vitorpamplona.quartz.marmot.mls.crypto
import com.vitorpamplona.quartz.utils.RandomInstance
/**
* Linux/Native X25519 implementation using pure Kotlin field arithmetic.
*
* Implements RFC 7748 X25519 Diffie-Hellman key agreement via Montgomery ladder.
* Key format: raw 32-byte Curve25519 keys (little-endian per RFC 7748).
*/
actual object X25519 {
actual fun generateKeyPair(): X25519KeyPair = TODO("X25519 not yet implemented for this platform")
private const val KEY_LENGTH = 32
actual fun generateKeyPair(): X25519KeyPair {
val privateKey = RandomInstance.bytes(KEY_LENGTH)
val publicKey = publicFromPrivate(privateKey)
return X25519KeyPair(privateKey, publicKey)
}
actual fun dh(
privateKey: ByteArray,
publicKey: ByteArray,
): ByteArray = TODO("X25519 not yet implemented for this platform")
): ByteArray {
require(privateKey.size == KEY_LENGTH) { "Private key must be 32 bytes" }
require(publicKey.size == KEY_LENGTH) { "Public key must be 32 bytes" }
actual fun publicFromPrivate(privateKey: ByteArray): ByteArray = TODO("X25519 not yet implemented for this platform")
val result = scalarmult(privateKey, publicKey)
require(!result.all { it == 0.toByte() }) {
"DH produced all-zero shared secret (possible small-subgroup attack)"
}
return result
}
actual fun publicFromPrivate(privateKey: ByteArray): ByteArray {
require(privateKey.size == KEY_LENGTH) { "Private key must be 32 bytes" }
val basepoint = ByteArray(KEY_LENGTH)
basepoint[0] = 9
return scalarmult(privateKey, basepoint)
}
/**
* X25519 scalar multiplication via Montgomery ladder (RFC 7748).
*/
private fun scalarmult(
n: ByteArray,
p: ByteArray,
): ByteArray {
val z = n.copyOf()
z[0] = (z[0].toInt() and 248).toByte()
z[31] = ((z[31].toInt() and 127) or 64).toByte()
val x = Curve25519Field.unpack25519(p)
val a = Curve25519Field.GF1.copyOf()
val b = x.copyOf()
val c = Curve25519Field.GF0.copyOf()
val d = Curve25519Field.GF1.copyOf()
for (i in 254 downTo 0) {
val r = ((z[i shr 3].toLong() shr (i and 7)) and 1)
Curve25519Field.sel25519(a, b, r)
Curve25519Field.sel25519(c, d, r)
val e = Curve25519Field.add(a, c)
val aMc = Curve25519Field.sub(a, c)
val f = Curve25519Field.add(b, d)
val bMd = Curve25519Field.sub(b, d)
val dd = Curve25519Field.sqr(e)
val ff = Curve25519Field.sqr(aMc)
val da = Curve25519Field.mul(bMd, e)
val cb = Curve25519Field.mul(f, aMc)
val ePrime = Curve25519Field.add(da, cb)
val aPrime = Curve25519Field.sub(da, cb)
val bNew = Curve25519Field.sqr(ePrime)
val aSqr = Curve25519Field.sqr(aPrime)
val dNew = Curve25519Field.mul(aSqr, x)
val aNew = Curve25519Field.mul(dd, ff)
val cc = Curve25519Field.sub(dd, ff)
val tmp = Curve25519Field.mul(cc, Curve25519Field.A24)
val ddPlusTmp = Curve25519Field.add(dd, tmp)
val cNew = Curve25519Field.mul(cc, ddPlusTmp)
aNew.copyInto(a)
bNew.copyInto(b)
cNew.copyInto(c)
dNew.copyInto(d)
Curve25519Field.sel25519(a, b, r)
Curve25519Field.sel25519(c, d, r)
}
val invC = Curve25519Field.inv25519(c)
val result = Curve25519Field.mul(a, invC)
return Curve25519Field.pack25519(result)
}
}
@@ -0,0 +1,267 @@
/*
* Copyright (c) 2025 Vitor Pamplona
*
* Permission is hereby granted, free of charge, to any person obtaining a copy of
* this software and associated documentation files (the "Software"), to deal in
* the Software without restriction, including without limitation the rights to use,
* copy, modify, merge, publish, distribute, sublicense, and/or sell copies of the
* Software, and to permit persons to whom the Software is furnished to do so,
* subject to the following conditions:
*
* The above copyright notice and this permission notice shall be included in all
* copies or substantial portions of the Software.
*
* THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR
* IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, FITNESS
* FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE AUTHORS OR
* COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER LIABILITY, WHETHER IN
* AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, OUT OF OR IN CONNECTION
* WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE SOFTWARE.
*/
package com.vitorpamplona.quartz.marmot.mls.crypto
/**
* Field arithmetic over GF(2^255-19) for Curve25519 operations.
*
* Field elements are represented as LongArray(16) in radix-2^16.
* Based on the TweetNaCl algorithm by Bernstein et al.
*/
internal object Curve25519Field {
/** The constant a24 = 121665, used in the Montgomery ladder. */
val A24 = gf(0xDB41L, 1)
/** 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)
/** Carry and reduce a field element. */
fun car25519(o: LongArray) {
for (i in 0 until 16) {
o[i] += (1L shl 16)
val c = o[i] shr 16
o[(i + 1) % 16] += c - 1 + (if (i == 15) 37 * (c - 1) else 0)
o[i] -= 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()
}
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
}
/** Field addition: o = a + b. */
fun add(
a: LongArray,
b: LongArray,
): LongArray {
val o = LongArray(16)
for (i in 0 until 16) o[i] = a[i] + b[i]
return o
}
/** Field subtraction: o = a - b. */
fun sub(
a: LongArray,
b: LongArray,
): LongArray {
val o = LongArray(16)
for (i in 0 until 16) o[i] = a[i] - b[i]
return o
}
/** Field multiplication: o = a * b (mod p). */
fun mul(
a: LongArray,
b: LongArray,
): LongArray {
val t = LongArray(31)
for (i in 0 until 16) {
for (j in 0 until 16) {
t[i + j] += a[i] * b[j]
}
}
for (i in 0 until 15) {
t[i] += 38 * t[i + 16]
}
val o = LongArray(16)
for (i in 0 until 16) o[i] = t[i]
car25519(o)
car25519(o)
return o
}
/** Field squaring: o = a^2 (mod p). */
fun sqr(a: LongArray): LongArray = mul(a, a)
/** Field inversion: o = a^(-1) (mod p) using Fermat's little theorem. */
fun inv25519(a: LongArray): LongArray {
var c = a.copyOf()
for (i in 253 downTo 0) {
c = sqr(c)
if (i != 2 && i != 4) c = mul(c, a)
}
return c
}
/** Parity of a field element (lowest bit after reduction). */
fun par25519(a: LongArray): Int {
val d = pack25519(a)
return d[0].toInt() and 1
}
/** Raise a field element to the power (2^252 - 3), used in sqrt. */
fun pow2523(a: LongArray): LongArray {
var c = a.copyOf()
for (i in 250 downTo 0) {
c = sqr(c)
if (i != 1) c = mul(c, a)
}
return c
}
}