From 9c10556e23e8ee445a97ea6b237e47058967f02a Mon Sep 17 00:00:00 2001 From: Ken MacKay Date: Tue, 14 May 2013 22:48:50 -0700 Subject: [PATCH] Removed reliance on memeset/memcpy, and added code for supporting processors with no efficient way of doing 32x32->64 bit multiplication (if ECC_MULT64 is set to 0, the code will only use 32x32->32 bit multiplication). This is for better support of Cortex M0 chips. --- ecdh.c | 167 +++++++++++++++++++++++++++++++++++++++++++++------------ ecdh.h | 1 + 2 files changed, 135 insertions(+), 33 deletions(-) diff --git a/ecdh.c b/ecdh.c index 6f11f3e..b189cf3 100644 --- a/ecdh.c +++ b/ecdh.c @@ -43,6 +43,15 @@ static uint32_t curve_p[NUM_ECC_DIGITS] = CONCAT(Curve_P_, ECC_CURVE); static EccPoint curve_G = CONCAT(Curve_G_, ECC_CURVE); static uint32_t curve_n[NUM_ECC_DIGITS] = CONCAT(Curve_N_, ECC_CURVE); +static void fast_clear(uint32_t *p_array) +{ + uint i; + for(i=0; i> 16; + uint32_t b0 = p_right[k-i] & 0xffff; + uint32_t b1 = p_right[k-i] >> 16; + + uint64_t l_product = (a0 * b0) + (((uint64_t)(a0 * b1) + a1 * b0) << 16) + (((uint64_t)(a1 * b1)) << 32); + r01 += l_product; + r2 += (r01 < l_product); + } + p_result[k] = (uint32_t)r01; + r01 = (r01 >> 32) | (((uint64_t)r2) << 32); + r2 = 0; + } + + p_result[NUM_ECC_DIGITS*2 - 1] = (uint32_t)r01; +} +#endif /* Computes p_result = (p_left + p_right) % p_mod. Assumes that p_left < p_mod and p_right < p_mod, p_result != p_mod. */ @@ -257,7 +299,7 @@ static void vli_mmod(uint32_t *p_result, uint32_t *p_product, uint32_t *p_mod) uint32_t l_tmp[4]; int l_carry; - memcpy(p_result, p_product, 4*sizeof(uint32_t)); + vli_set(p_result, p_product); l_tmp[0] = p_product[4]; l_tmp[1] = p_product[5]; @@ -310,7 +352,7 @@ static void vli_mmod(uint32_t *p_result, uint32_t *p_product, uint32_t *p_mod) uint32_t l_tmp[5]; int l_carry; - memcpy(p_result, p_product, 5*sizeof(uint32_t)); + vli_set(p_result, p_product); l_tmp[0] = (p_product[5] & 0x7FFFFFFF) | (p_product[5] << 31); l_tmp[1] = (p_product[5] >> 1) | (p_product[6] << 31); @@ -348,13 +390,16 @@ static void vli_mmod(uint32_t *p_result, uint32_t *p_product, uint32_t *p_mod) uint32_t l_tmp[6]; int l_carry; - memcpy(p_result, p_product, 6*sizeof(uint32_t)); + vli_set(p_result, p_product); - memcpy(l_tmp, &p_product[6], 6*sizeof(uint32_t)); + vli_set(l_tmp, &p_product[6]); l_carry = vli_add(p_result, p_result, l_tmp); l_tmp[0] = l_tmp[1] = 0; - memcpy(&l_tmp[2], &p_product[6], 4*sizeof(uint32_t)); + l_tmp[2] = p_product[6]; + l_tmp[3] = p_product[7]; + l_tmp[4] = p_product[8]; + l_tmp[5] = p_product[9]; l_carry += vli_add(p_result, p_result, l_tmp); l_tmp[0] = l_tmp[2] = p_product[10]; @@ -377,21 +422,28 @@ static void vli_mmod(uint32_t *p_result, uint32_t *p_product, uint32_t *p_mod) uint32_t l_tmp[7]; int l_carry; - memcpy(p_result, p_product, 7*sizeof(uint32_t)); + vli_set(p_result, p_product); l_tmp[0] = l_tmp[1] = l_tmp[2] = 0; - memcpy(&l_tmp[3], &p_product[7], 4*sizeof(uint32_t)); + l_tmp[3] = p_product[7]; + l_tmp[4] = p_product[8]; + l_tmp[5] = p_product[9]; + l_tmp[6] = p_product[10]; l_carry = vli_add(p_result, p_result, l_tmp); + l_tmp[3] = p_product[11]; + l_tmp[4] = p_product[12]; + l_tmp[5] = p_product[13]; l_tmp[6] = 0; - memcpy(&l_tmp[3], &p_product[11], 3*sizeof(uint32_t)); l_carry += vli_add(p_result, p_result, l_tmp); - memcpy(l_tmp, &p_product[7], 7*sizeof(uint32_t)); + vli_set(l_tmp, &p_product[7]); l_carry -= vli_sub(p_result, p_result, l_tmp); + l_tmp[0] = p_product[11]; + l_tmp[1] = p_product[12]; + l_tmp[2] = p_product[13]; l_tmp[3] = l_tmp[4] = l_tmp[5] = l_tmp[6] = 0; - memcpy(l_tmp, &p_product[11], 3*sizeof(uint32_t)); l_carry -= vli_sub(p_result, p_result, l_tmp); if(l_carry < 0) @@ -420,16 +472,23 @@ static void vli_mmod(uint32_t *p_result, uint32_t *p_product, uint32_t *p_mod) int l_carry; /* t */ - memcpy(p_result, p_product, 8*sizeof(uint32_t)); + vli_set(p_result, p_product); /* s1 */ l_tmp[0] = l_tmp[1] = l_tmp[2] = 0; - memcpy(&l_tmp[3], &p_product[11], 5*sizeof(uint32_t)); + l_tmp[3] = p_product[11]; + l_tmp[4] = p_product[12]; + l_tmp[5] = p_product[13]; + l_tmp[6] = p_product[14]; + l_tmp[7] = p_product[15]; l_carry = vli_lshift(l_tmp, l_tmp, 1); l_carry += vli_add(p_result, p_result, l_tmp); /* s2 */ - memcpy(&l_tmp[3], &p_product[12], 4*sizeof(uint32_t)); + l_tmp[3] = p_product[12]; + l_tmp[4] = p_product[13]; + l_tmp[5] = p_product[14]; + l_tmp[6] = p_product[15]; l_tmp[7] = 0; l_carry += vli_lshift(l_tmp, l_tmp, 1); l_carry += vli_add(p_result, p_result, l_tmp); @@ -524,6 +583,7 @@ static void vli_modMult(uint32_t *p_result, uint32_t *p_left, uint32_t *p_right, #if ECC_SQUARE_FUNC /* Computes p_result = p_left^2. */ +#if ECC_MULT64 static void vli_square(uint32_t *p_result, uint32_t *p_left) { uint64_t r01 = 0; @@ -538,7 +598,7 @@ static void vli_square(uint32_t *p_result, uint32_t *p_left) uint64_t l_product = (uint64_t)p_left[i] * p_left[k-i]; if(i < k-i) { - r2 += !!(l_product & ((uint64_t)1 << 63)); + r2 += l_product >> 63; l_product *= 2; } r01 += l_product; @@ -551,6 +611,41 @@ static void vli_square(uint32_t *p_result, uint32_t *p_left) p_result[NUM_ECC_DIGITS*2 - 1] = (uint32_t)r01; } +#else +static void vli_square(uint32_t *p_result, uint32_t *p_left) +{ + uint64_t r01 = 0; + uint32_t r2 = 0; + + uint i, k; + for(k=0; k < NUM_ECC_DIGITS*2 - 1; ++k) + { + uint l_min = (k < NUM_ECC_DIGITS ? 0 : (k + 1) - NUM_ECC_DIGITS); + for(i=l_min; i<=k && i<=k-i; ++i) + { + uint32_t a0 = p_left[i] & 0xffff; + uint32_t a1 = p_left[i] >> 16; + uint32_t b0 = p_left[k-i] & 0xffff; + uint32_t b1 = p_left[k-i] >> 16; + + uint64_t l_product = (a0 * b0) + (((uint64_t)(a0 * b1) + a1 * b0) << 16) + (((uint64_t)(a1 * b1)) << 32); + + if(i < k-i) + { + r2 += l_product >> 63; + l_product *= 2; + } + r01 += l_product; + r2 += (r01 < l_product); + } + p_result[k] = (uint32_t)r01; + r01 = (r01 >> 32) | (((uint64_t)r2) << 32); + r2 = 0; + } + + p_result[NUM_ECC_DIGITS*2 - 1] = (uint32_t)r01; +} +#endif /* Computes p_result = p_left^2 % p_mod. */ static void vli_modSquare(uint32_t *p_result, uint32_t *p_left, uint32_t *p_mod) @@ -579,7 +674,7 @@ static void vli_modDiv(uint32_t *p_result, uint32_t *p_left, uint32_t *p_right, vli_set(a, p_right); vli_set(b, p_mod); vli_set(u, p_left); - memset(v, 0, NUM_ECC_DIGITS*sizeof(uint32_t)); + fast_clear(v); while ((l_cmpResult = vli_cmp(a, b)) != 0) { @@ -656,9 +751,10 @@ static void vli_modDiv(uint32_t *p_result, uint32_t *p_left, uint32_t *p_right, /* Computes p_result = (1 / p_input) % p_mod. All VLIs are the same size. */ static void vli_modInv(uint32_t *p_result, uint32_t *p_input, uint32_t *p_mod) { - uint32_t n[NUM_ECC_DIGITS]; - memset(n, 0, NUM_ECC_DIGITS*sizeof(uint32_t)); - n[0] = 1; + uint32_t n[NUM_ECC_DIGITS]; + + fast_clear(n); + n[0] = 1; vli_modDiv(p_result, n, p_input, p_mod); } @@ -668,8 +764,8 @@ static void vli_modInv(uint32_t *p_result, uint32_t *p_input, uint32_t *p_mod) /* Clears a point (set it to the point at infinity). */ static void EccPoint_clear(EccPoint *p_point) { - memset(p_point->x, 0, NUM_ECC_DIGITS*sizeof(uint32_t)); - memset(p_point->y, 0, NUM_ECC_DIGITS*sizeof(uint32_t)); + fast_clear(p_point->x); + fast_clear(p_point->y); } /* Copies a point. */ @@ -700,7 +796,7 @@ static void EccPoint_double_projective(EccPoint *P3, uint32_t *Z3, EccPoint *P1, if(vli_zero(Z1)) { - memset(Z3, 0, sizeof(uint32_t) * NUM_ECC_DIGITS); + fast_clear(Z3); return; } @@ -756,7 +852,7 @@ static void EccPoint_add_mixed(EccPoint *P3, uint32_t *Z3, EccPoint *P1, uint32_ if(vli_zero(Z1)) { EccPoint_copy(P3, P2); - memset(Z3, 0, sizeof(uint32_t) * NUM_ECC_DIGITS); + fast_clear(Z3); Z3[0] = 1; return; } @@ -779,7 +875,7 @@ static void EccPoint_add_mixed(EccPoint *P3, uint32_t *Z3, EccPoint *P1, uint32_ } else { - memset(Z3, 0, sizeof(uint32_t) * NUM_ECC_DIGITS); + fast_clear(Z3); } return; } @@ -810,17 +906,21 @@ static void EccPoint_add_mixed(EccPoint *P3, uint32_t *Z3, EccPoint *P1, uint32_ void EccPoint_mult(EccPoint *p_result, EccPoint *p_point, uint32_t *p_scalar) { uint32_t l_tmp[NUM_ECC_DIGITS]; - uint32_t Z1[NUM_ECC_DIGITS] = {0}; - uint32_t l_plus[NUM_ECC_DIGITS] = {0}; - uint32_t l_minus[NUM_ECC_DIGITS] = {0}; + uint32_t Z1[NUM_ECC_DIGITS]; + uint32_t l_plus[NUM_ECC_DIGITS]; + uint32_t l_minus[NUM_ECC_DIGITS]; int l_numBits = vli_numBits(p_scalar); uint l_carry; int i; + fast_clear(Z1); + fast_clear(l_plus); + fast_clear(l_minus); + EccPoint l_neg; - memcpy(l_neg.x, p_point->x, NUM_ECC_DIGITS*sizeof(uint32_t)); + vli_set(l_neg.x, p_point->x); vli_sub(l_neg.y, curve_p, p_point->y); l_carry = 0; @@ -881,11 +981,12 @@ void EccPoint_mult(EccPoint *p_result, EccPoint *p_point, uint32_t *p_scalar) void EccPoint_mult(EccPoint *p_result, EccPoint *p_point, uint32_t *p_scalar) { uint32_t l_tmp[NUM_ECC_DIGITS]; - uint32_t Z1[NUM_ECC_DIGITS] = {0}; + uint32_t Z1[NUM_ECC_DIGITS]; uint l_numBits = vli_numBits(p_scalar); int i; + fast_clear(Z1); EccPoint_clear(p_result); for(i = l_numBits - 1; i >= 0; --i) @@ -917,7 +1018,7 @@ int ecdh_shared_secret(uint32_t p_secret[NUM_ECC_DIGITS], EccPoint *p_publicKey, return 0; } - memcpy(p_secret, l_product.x, NUM_ECC_DIGITS * sizeof(uint32_t)); + vli_set(p_secret, l_product.x); return 1; } diff --git a/ecdh.h b/ecdh.h index 396b00d..7461914 100644 --- a/ecdh.h +++ b/ecdh.h @@ -11,6 +11,7 @@ ECC_USE_NAF - If enabled, this will convert the private key to a non-adjacent fo */ #define ECC_SQUARE_FUNC 1 #define ECC_USE_NAF 1 +#define ECC_MULT64 1 #define ECC_CURVE secp160r1