src/crypto/curve25519/field.bend source
src/crypto/curve25519/field.bend on the hub · documented module
import Base# Arithmetic in GF(p), p = 2^255 - 19, for X25519 and Ed25519.## A field element is a list of 32 U32 limbs, little-endian in radix 2^8:# value(a) = a0 + 2^8 a1 + ... + 2^248 a31 (spec/crypto/curve25519/field.bend).# Every operation returns a "tight" element: 32 limbs, each below 2^8, so# its value is below 2^256 (not necessarily below p; `freeze` gives the# canonical representative, which is also its 32-byte encoding).## Only U32 addition, multiplication, subtraction and division/remainder by# the constant 256 (a shift and a mask after compilation) are used, and no# intermediate reaches 2^32 (proved: proofs/crypto/curve25519). The schoolbook# product of two tight elements has limbs below 32 * 255^2 < 2^21; folding# the high half by 2^256 == 38 (mod p) keeps them below 2^27; three carry# passes bring the result back to tight (TweetNaCl's car25519, in radix# 2^8; Fiat-Crypto's unsaturated-limb reduction, Erbsen et al. 2019).## No function branches on a limb value: selection is arithmetic# (a * (1 - s) + b * s), list shapes are fixed (32 limbs), and exponents# are public. Bend has no timing model, so constant time is by# construction, not proved.# ---- limb lists ----# limbwise sum; the longer tail is keptdef addl(xs: List<&2, U32>, ys: List<&2, U32>) -> List<&2, U32>: match xs: case Nil{}: ys case Con{x, xt}: match ys: case Nil{}: Con{x, xt} case Con{y, yt}: Con{U32.add(x, y), addl(xt, yt)}# every limb times adef scal(+a: U32, ys: List<&2, U32>) -> List<&2, U32>: match ys: case Nil{}: Nil{} case Con{y, yt}: Con{U32.mul(a, y), scal(a, yt)}# the product polynomial: conv(x :: xs, ys) = x ys + 2^8 conv(xs, ys)def conv(xs: List<&2, U32>, +ys: List<&2, U32>) -> List<&2, U32>: match xs: case Nil{}: Nil{} case Con{x, xt}: addl(scal(x, ys), Con{0, conv(xt, ys)})def take(n: Nat, xs: List<&2, U32>) -> List<&2, U32>: match n: case 0n: Nil{} case 1n+m: match xs: case Nil{}: Nil{} case Con{x, xt}: Con{x, take(m, xt)}def drop(n: Nat, xs: List<&2, U32>) -> List<&2, U32>: match n: case 0n: xs case 1n+m: match xs: case Nil{}: Nil{} case Con{x, xt}: drop(m, xt)# ---- carries ----# limbs and the carry out of the top onetype Cr is Data: Cr{limbs: List<&2, U32>, out: U32}def cr_limbs(r: Cr) -> List<&2, U32>: match r: case Cr{a, c}: adef cr_out(r: Cr) -> U32: match r: case Cr{a, c}: cdef carry_con(l: U32, r: Cr) -> Cr: match r: case Cr{a, c}: Cr{Con{l, a}, c}# one carry pass: limbs below 2^8 and the carry out of the top limbdef carry(xs: List<&2, U32>, c: U32) -> Cr: match xs: case Nil{}: Cr{Nil{}, c} case Con{h, t}: +s = U32.add(h, c) carry_con(U32.mod(s, 256), carry(t, U32.div(s, 256)))# add k to the lowest limbdef add0(xs: List<&2, U32>, k: U32) -> List<&2, U32>: match xs: case Nil{}: Nil{} case Con{x, xt}: Con{U32.add(x, k), xt}# carry, then fold the carry out back in: 2^256 == 38 (mod p)def pass_fin(r: Cr) -> List<&2, U32>: match r: case Cr{a, c}: add0(a, U32.mul(38, c))def pass(xs: List<&2, U32>) -> List<&2, U32>: pass_fin(carry(xs, 0))# 32 limbs below 2^27 to a tight element of the same value mod pdef reduce(xs: List<&2, U32>) -> List<&2, U32>: pass(pass(pass(xs)))# ---- constants ----def zero() -> List<&2, U32>: [0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0]def one() -> List<&2, U32>: [1, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0]# a small constant k < 2^8 as an elementdef small(k: U32) -> List<&2, U32>: [k, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0]# 8p limbwise: every limb at least 2^10 - 8, so a + 8p - b never borrowsdef eight_p() -> List<&2, U32>: [1896, 2040, 2040, 2040, 2040, 2040, 2040, 2040, 2040, 2040, 2040, 2040, 2040, 2040, 2040, 2040, 2040, 2040, 2040, 2040, 2040, 2040, 2040, 2040, 2040, 2040, 2040, 2040, 2040, 2040, 2040, 1016]# 2^256 - p = 2^255 + 19def comp_p() -> List<&2, U32>: [19, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 128]# ---- field operations (tight in, tight out) ----# a + k - b limbwisedef subl(xs: List<&2, U32>, ys: List<&2, U32>, ks: List<&2, U32>) -> List<&2, U32>: match xs: case Nil{}: Nil{} case Con{x, xt}: match ys: case Nil{}: Nil{} case Con{y, yt}: match ks: case Nil{}: Nil{} case Con{k, kt}: Con{U32.sub(U32.add(x, k), y), subl(xt, yt, kt)}def add(a: List<&2, U32>, b: List<&2, U32>) -> List<&2, U32>: reduce(addl(a, b))def sub(a: List<&2, U32>, b: List<&2, U32>) -> List<&2, U32>: reduce(subl(a, b, eight_p()))# fold the product: low 32 limbs + 38 * high limbsdef wide(+zs: List<&2, U32>) -> List<&2, U32>: addl(take(32n, zs), scal(38, drop(32n, zs)))def mul(a: List<&2, U32>, +b: List<&2, U32>) -> List<&2, U32>: reduce(wide(conv(a, b)))def sq(+a: List<&2, U32>) -> List<&2, U32>: mul(a, a)# a * k for a constant k < 2^17def mul_small(a: List<&2, U32>, +k: U32) -> List<&2, U32>: reduce(scal(k, a))def neg(a: List<&2, U32>) -> List<&2, U32>: sub(zero(), a)# ---- selection (s is 0 or 1) ----# s == 0: a; s == 1: b; limbwise a * (1 - s) + b * sdef select(+s: U32, xs: List<&2, U32>, ys: List<&2, U32>) -> List<&2, U32>: match xs: case Nil{}: Nil{} case Con{x, xt}: match ys: case Nil{}: Nil{} case Con{y, yt}: Con{U32.add(U32.mul(x, U32.sub(1, s)), U32.mul(y, s)), select(s, xt, yt)}# ---- powers (public exponents) ----## The loops take the base x first and inspect it before their counter, so# the proof checker never unfolds a fixed-count loop over an unknown base# (it keeps pow_ones(x, 249n, x) as a call until x is known).# acc^(2^n) * x^(2^n - 1): n steps of acc = acc^2 * xdef pow_ones(+x: List<&2, U32>, n: Nat, acc: List<&2, U32>) -> List<&2, U32>: match x: case Nil{}: acc case Con{h, t}: match n: case 0n: acc case 1n+m: pow_ones(x, m, mul(sq(acc), x))# one bit of a public exponent, most significant first: acc^2 (* x)def pow_bit(b: Bool, +x: List<&2, U32>, acc: List<&2, U32>) -> List<&2, U32>: match b: case True{}: mul(sq(acc), x) case False{}: sq(acc)def pow_bits(+x: List<&2, U32>, bs: List<&2, Bool>, acc: List<&2, U32>) -> List<&2, U32>: match x: case Nil{}: acc case Con{h, t}: match bs: case Nil{}: acc case Con{b, bt}: pow_bits(x, bt, pow_bit(b, x, acc))# x^(p - 2) = x^(2^255 - 21): 250 one bits (x, then 249 steps), then 0 1 0 1 1def inv(+x: List<&2, U32>) -> List<&2, U32>: pow_bits(x, [False{}, True{}, False{}, True{}, True{}], pow_ones(x, 249n, x))# x^((p - 5) / 8) = x^(2^252 - 3): 250 one bits, then 0 1def pow_p58(+x: List<&2, U32>) -> List<&2, U32>: pow_bits(x, [False{}, True{}], pow_ones(x, 249n, x))# ---- canonical form ----def csub_fin(r: Cr, x: List<&2, U32>) -> List<&2, U32>: match r: case Cr{s, c}: select(c, x, s)# x - p when x >= p, else x (x tight): x + 2^256 - p carries out iff x >= pdef csub(+x: List<&2, U32>) -> List<&2, U32>: csub_fin(carry(addl(x, comp_p()), 0), x)# the canonical representative, below p: 2^256 < 3pdef freeze(+x: List<&2, U32>) -> List<&2, U32>: csub(csub(x))# ---- bytes ----# the 32-byte little-endian encoding of a canonical element is its limbsdef to_bytes(+x: List<&2, U32>) -> List<&2, U32>: freeze(x)# the last byte with its top bit cleared (RFC 7748 decodeUCoordinate)def mask_top(xs: List<&2, U32>) -> List<&2, U32>: match xs: case Nil{}: Nil{} case Con{x, xt}: match xt: case Nil{}: Con{U32.and(x, 127), Nil{}} case Con{y, yt}: Con{x, mask_top(Con{y, yt})}# bytes to an element, bit 255 ignoreddef of_bytes(bs: List<&2, U32>) -> List<&2, U32>: mask_top(bs)# the sum of the limbs (below 2^13 for 32 bytes); no early exitdef sum_all(xs: List<&2, U32>, acc: U32) -> U32: match xs: case Nil{}: acc case Con{x, xt}: sum_all(xt, U32.add(acc, x))# x == 0 in the field: every byte of the canonical form is 0def is_zero(+x: List<&2, U32>) -> Bool: U32.is_eq(sum_all(freeze(x), 0), 0)# a == b in the fielddef eq(+a: List<&2, U32>, +b: List<&2, U32>) -> Bool: is_zero(sub(a, b))# the parity of the canonical representative (RFC 8032 x_0)def low_bit(xs: List<&2, U32>) -> U32: match xs: case Nil{}: 0 case Con{l, lt}: U32.and(l, 1)def parity(+x: List<&2, U32>) -> U32: low_bit(freeze(x))