srp.c 14 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329330331332333334335336337338339340341342343344345346347348349350351352353354355356357358359360361362363364365366367368369370371372373374375376377378379380381382383384385386387388389390391392393394395396397398399400401402403404405406407408409410411412413414415416417418419420421422423424425426427428429430431432433434435436437438439440441442443444445
  1. #include <string.h>
  2. #include "esp_log.h"
  3. #include "esp_random.h"
  4. #include "mbedtls/bignum.h"
  5. #include "sodium.h"
  6. #include "srp.h"
  7. static const char *TAG = "srp";
  8. // SRP-6a 3072-bit prime N (from RFC 5054)
  9. static const uint8_t srp_N[] = {
  10. 0xFF, 0xFF, 0xFF, 0xFF, 0xFF, 0xFF, 0xFF, 0xFF, 0xC9, 0x0F, 0xDA, 0xA2,
  11. 0x21, 0x68, 0xC2, 0x34, 0xC4, 0xC6, 0x62, 0x8B, 0x80, 0xDC, 0x1C, 0xD1,
  12. 0x29, 0x02, 0x4E, 0x08, 0x8A, 0x67, 0xCC, 0x74, 0x02, 0x0B, 0xBE, 0xA6,
  13. 0x3B, 0x13, 0x9B, 0x22, 0x51, 0x4A, 0x08, 0x79, 0x8E, 0x34, 0x04, 0xDD,
  14. 0xEF, 0x95, 0x19, 0xB3, 0xCD, 0x3A, 0x43, 0x1B, 0x30, 0x2B, 0x0A, 0x6D,
  15. 0xF2, 0x5F, 0x14, 0x37, 0x4F, 0xE1, 0x35, 0x6D, 0x6D, 0x51, 0xC2, 0x45,
  16. 0xE4, 0x85, 0xB5, 0x76, 0x62, 0x5E, 0x7E, 0xC6, 0xF4, 0x4C, 0x42, 0xE9,
  17. 0xA6, 0x37, 0xED, 0x6B, 0x0B, 0xFF, 0x5C, 0xB6, 0xF4, 0x06, 0xB7, 0xED,
  18. 0xEE, 0x38, 0x6B, 0xFB, 0x5A, 0x89, 0x9F, 0xA5, 0xAE, 0x9F, 0x24, 0x11,
  19. 0x7C, 0x4B, 0x1F, 0xE6, 0x49, 0x28, 0x66, 0x51, 0xEC, 0xE4, 0x5B, 0x3D,
  20. 0xC2, 0x00, 0x7C, 0xB8, 0xA1, 0x63, 0xBF, 0x05, 0x98, 0xDA, 0x48, 0x36,
  21. 0x1C, 0x55, 0xD3, 0x9A, 0x69, 0x16, 0x3F, 0xA8, 0xFD, 0x24, 0xCF, 0x5F,
  22. 0x83, 0x65, 0x5D, 0x23, 0xDC, 0xA3, 0xAD, 0x96, 0x1C, 0x62, 0xF3, 0x56,
  23. 0x20, 0x85, 0x52, 0xBB, 0x9E, 0xD5, 0x29, 0x07, 0x70, 0x96, 0x96, 0x6D,
  24. 0x67, 0x0C, 0x35, 0x4E, 0x4A, 0xBC, 0x98, 0x04, 0xF1, 0x74, 0x6C, 0x08,
  25. 0xCA, 0x18, 0x21, 0x7C, 0x32, 0x90, 0x5E, 0x46, 0x2E, 0x36, 0xCE, 0x3B,
  26. 0xE3, 0x9E, 0x77, 0x2C, 0x18, 0x0E, 0x86, 0x03, 0x9B, 0x27, 0x83, 0xA2,
  27. 0xEC, 0x07, 0xA2, 0x8F, 0xB5, 0xC5, 0x5D, 0xF0, 0x6F, 0x4C, 0x52, 0xC9,
  28. 0xDE, 0x2B, 0xCB, 0xF6, 0x95, 0x58, 0x17, 0x18, 0x39, 0x95, 0x49, 0x7C,
  29. 0xEA, 0x95, 0x6A, 0xE5, 0x15, 0xD2, 0x26, 0x18, 0x98, 0xFA, 0x05, 0x10,
  30. 0x15, 0x72, 0x8E, 0x5A, 0x8A, 0xAA, 0xC4, 0x2D, 0xAD, 0x33, 0x17, 0x0D,
  31. 0x04, 0x50, 0x7A, 0x33, 0xA8, 0x55, 0x21, 0xAB, 0xDF, 0x1C, 0xBA, 0x64,
  32. 0xEC, 0xFB, 0x85, 0x04, 0x58, 0xDB, 0xEF, 0x0A, 0x8A, 0xEA, 0x71, 0x57,
  33. 0x5D, 0x06, 0x0C, 0x7D, 0xB3, 0x97, 0x0F, 0x85, 0xA6, 0xE1, 0xE4, 0xC7,
  34. 0xAB, 0xF5, 0xAE, 0x8C, 0xDB, 0x09, 0x33, 0xD7, 0x1E, 0x8C, 0x94, 0xE0,
  35. 0x4A, 0x25, 0x61, 0x9D, 0xCE, 0xE3, 0xD2, 0x26, 0x1A, 0xD2, 0xEE, 0x6B,
  36. 0xF1, 0x2F, 0xFA, 0x06, 0xD9, 0x8A, 0x08, 0x64, 0xD8, 0x76, 0x02, 0x73,
  37. 0x3E, 0xC8, 0x6A, 0x64, 0x52, 0x1F, 0x2B, 0x18, 0x17, 0x7B, 0x20, 0x0C,
  38. 0xBB, 0xE1, 0x17, 0x57, 0x7A, 0x61, 0x5D, 0x6C, 0x77, 0x09, 0x88, 0xC0,
  39. 0xBA, 0xD9, 0x46, 0xE2, 0x08, 0xE2, 0x4F, 0xA0, 0x74, 0xE5, 0xAB, 0x31,
  40. 0x43, 0xDB, 0x5B, 0xFC, 0xE0, 0xFD, 0x10, 0x8E, 0x4B, 0x82, 0xD1, 0x20,
  41. 0xA9, 0x3A, 0xD2, 0xCA, 0xFF, 0xFF, 0xFF, 0xFF, 0xFF, 0xFF, 0xFF, 0xFF};
  42. #define SRP_GENERATOR 5
  43. // Helper: write MPI to buffer with minimum bytes (no leading zeros except for
  44. // value 0)
  45. static size_t mpi_to_bytes_min(const mbedtls_mpi *mpi, uint8_t *buf,
  46. size_t len) {
  47. size_t mpi_size = mbedtls_mpi_size(mpi);
  48. if (mpi_size == 0) {
  49. if (len < 1) {
  50. return 0;
  51. }
  52. buf[0] = 0;
  53. return 1;
  54. }
  55. if (mpi_size > len) {
  56. return 0;
  57. }
  58. if (mbedtls_mpi_write_binary(mpi, buf, mpi_size) != 0) {
  59. return 0;
  60. }
  61. return mpi_size;
  62. }
  63. // Helper: write MPI to buffer, zero-padded to fixed length
  64. static int mpi_to_bytes_padded(const mbedtls_mpi *mpi, uint8_t *buf,
  65. size_t len) {
  66. size_t mpi_size = mbedtls_mpi_size(mpi);
  67. if (mpi_size > len) {
  68. return -1;
  69. }
  70. memset(buf, 0, len);
  71. return mbedtls_mpi_write_binary(mpi, buf + (len - mpi_size), mpi_size);
  72. }
  73. // Helper: trim leading zeros from buffer
  74. static void trim_leading_zeros(const uint8_t *in, size_t in_len,
  75. const uint8_t **out, size_t *out_len) {
  76. while (in_len > 1 && *in == 0) {
  77. in++;
  78. in_len--;
  79. }
  80. *out = in;
  81. *out_len = in_len;
  82. }
  83. // Compute M1 = H(H(N)^H(g) || H(I) || s || A || B || K)
  84. static void compute_m1(uint8_t *out, const uint8_t *h_Ng_xor,
  85. const uint8_t *h_I, const uint8_t *salt, size_t salt_len,
  86. const uint8_t *A, size_t A_len, const uint8_t *B,
  87. size_t B_len, const uint8_t *K, size_t K_len) {
  88. crypto_hash_sha512_state state;
  89. crypto_hash_sha512_init(&state);
  90. crypto_hash_sha512_update(&state, h_Ng_xor, 64);
  91. crypto_hash_sha512_update(&state, h_I, 64);
  92. crypto_hash_sha512_update(&state, salt, salt_len);
  93. crypto_hash_sha512_update(&state, A, A_len);
  94. crypto_hash_sha512_update(&state, B, B_len);
  95. crypto_hash_sha512_update(&state, K, K_len);
  96. crypto_hash_sha512_final(&state, out);
  97. }
  98. srp_session_t *srp_session_create(void) {
  99. srp_session_t *session = calloc(1, sizeof(srp_session_t));
  100. return session;
  101. }
  102. void srp_session_free(srp_session_t *session) {
  103. if (session) {
  104. memset(session, 0, sizeof(srp_session_t));
  105. free(session);
  106. }
  107. }
  108. esp_err_t srp_start(srp_session_t *session, const char *username,
  109. const char *password) {
  110. if (!session || !username || !password) {
  111. return ESP_ERR_INVALID_ARG;
  112. }
  113. mbedtls_mpi N, g, k, v, b, B, x, tmp, tmp2;
  114. mbedtls_mpi_init(&N);
  115. mbedtls_mpi_init(&g);
  116. mbedtls_mpi_init(&k);
  117. mbedtls_mpi_init(&v);
  118. mbedtls_mpi_init(&b);
  119. mbedtls_mpi_init(&B);
  120. mbedtls_mpi_init(&x);
  121. mbedtls_mpi_init(&tmp);
  122. mbedtls_mpi_init(&tmp2);
  123. int ret = -1;
  124. // Generate random salt
  125. esp_fill_random(session->salt, SRP_SALT_BYTES);
  126. // Load N and g
  127. mbedtls_mpi_read_binary(&N, srp_N, sizeof(srp_N));
  128. mbedtls_mpi_lset(&g, SRP_GENERATOR);
  129. // k = H(N || pad(g))
  130. {
  131. uint8_t hash_input[SRP_PRIME_BYTES * 2];
  132. memcpy(hash_input, srp_N, SRP_PRIME_BYTES);
  133. memset(hash_input + SRP_PRIME_BYTES, 0, SRP_PRIME_BYTES);
  134. hash_input[SRP_PRIME_BYTES * 2 - 1] = SRP_GENERATOR;
  135. uint8_t k_hash[64];
  136. crypto_hash_sha512(k_hash, hash_input, sizeof(hash_input));
  137. mbedtls_mpi_read_binary(&k, k_hash, 64);
  138. mbedtls_mpi_mod_mpi(&k, &k, &N);
  139. }
  140. // x = H(s || H(I || ":" || P))
  141. {
  142. uint8_t inner_hash[64];
  143. crypto_hash_sha512_state state;
  144. crypto_hash_sha512_init(&state);
  145. crypto_hash_sha512_update(&state, (const uint8_t *)username,
  146. strlen(username));
  147. crypto_hash_sha512_update(&state, (const uint8_t *)":", 1);
  148. crypto_hash_sha512_update(&state, (const uint8_t *)password,
  149. strlen(password));
  150. crypto_hash_sha512_final(&state, inner_hash);
  151. uint8_t x_hash[64];
  152. crypto_hash_sha512_init(&state);
  153. crypto_hash_sha512_update(&state, session->salt, SRP_SALT_BYTES);
  154. crypto_hash_sha512_update(&state, inner_hash, 64);
  155. crypto_hash_sha512_final(&state, x_hash);
  156. mbedtls_mpi_read_binary(&x, x_hash, 64);
  157. }
  158. // v = g^x mod N
  159. if (mbedtls_mpi_exp_mod(&v, &g, &x, &N, NULL) != 0) {
  160. goto cleanup;
  161. }
  162. // Generate random b (server secret)
  163. {
  164. uint8_t b_bytes[SRP_PRIME_BYTES];
  165. esp_fill_random(b_bytes, sizeof(b_bytes));
  166. mbedtls_mpi_read_binary(&b, b_bytes, sizeof(b_bytes));
  167. mbedtls_mpi_mod_mpi(&b, &b, &N);
  168. mpi_to_bytes_padded(&b, session->server_secret, SRP_PRIME_BYTES);
  169. }
  170. // B = (k*v + g^b) mod N
  171. if (mbedtls_mpi_exp_mod(&tmp, &g, &b, &N, NULL) != 0) {
  172. goto cleanup;
  173. }
  174. if (mbedtls_mpi_mul_mpi(&tmp2, &k, &v) != 0) {
  175. goto cleanup;
  176. }
  177. if (mbedtls_mpi_add_mpi(&B, &tmp2, &tmp) != 0) {
  178. goto cleanup;
  179. }
  180. mbedtls_mpi_mod_mpi(&B, &B, &N);
  181. mpi_to_bytes_padded(&B, session->server_public_key, SRP_PRIME_BYTES);
  182. session->state = 1;
  183. ret = 0;
  184. cleanup:
  185. mbedtls_mpi_free(&N);
  186. mbedtls_mpi_free(&g);
  187. mbedtls_mpi_free(&k);
  188. mbedtls_mpi_free(&v);
  189. mbedtls_mpi_free(&b);
  190. mbedtls_mpi_free(&B);
  191. mbedtls_mpi_free(&x);
  192. mbedtls_mpi_free(&tmp);
  193. mbedtls_mpi_free(&tmp2);
  194. return ret == 0 ? ESP_OK : ESP_FAIL;
  195. }
  196. const uint8_t *srp_get_salt(srp_session_t *session) {
  197. return session ? session->salt : NULL;
  198. }
  199. const uint8_t *srp_get_public_key(srp_session_t *session, size_t *len) {
  200. if (!session) {
  201. return NULL;
  202. }
  203. if (len) {
  204. *len = SRP_PRIME_BYTES;
  205. }
  206. return session->server_public_key;
  207. }
  208. esp_err_t srp_verify_client(srp_session_t *session,
  209. const uint8_t *client_public_key,
  210. size_t client_pk_len, const uint8_t *client_proof,
  211. size_t proof_len) {
  212. if (!session || !client_public_key || !client_proof ||
  213. proof_len < SRP_PROOF_BYTES) {
  214. return ESP_ERR_INVALID_ARG;
  215. }
  216. // Store client's public key A (zero-padded)
  217. if (client_pk_len > SRP_PRIME_BYTES) {
  218. client_pk_len = SRP_PRIME_BYTES;
  219. }
  220. memset(session->client_public_key, 0, SRP_PRIME_BYTES);
  221. memcpy(session->client_public_key + (SRP_PRIME_BYTES - client_pk_len),
  222. client_public_key, client_pk_len);
  223. mbedtls_mpi N, g, A, B, b, u, S, k, v, x, tmp, tmp2;
  224. mbedtls_mpi_init(&N);
  225. mbedtls_mpi_init(&g);
  226. mbedtls_mpi_init(&A);
  227. mbedtls_mpi_init(&B);
  228. mbedtls_mpi_init(&b);
  229. mbedtls_mpi_init(&u);
  230. mbedtls_mpi_init(&S);
  231. mbedtls_mpi_init(&k);
  232. mbedtls_mpi_init(&v);
  233. mbedtls_mpi_init(&x);
  234. mbedtls_mpi_init(&tmp);
  235. mbedtls_mpi_init(&tmp2);
  236. int ret = -1;
  237. // Load parameters
  238. mbedtls_mpi_read_binary(&N, srp_N, sizeof(srp_N));
  239. mbedtls_mpi_lset(&g, SRP_GENERATOR);
  240. mbedtls_mpi_read_binary(&A, session->client_public_key, SRP_PRIME_BYTES);
  241. mbedtls_mpi_read_binary(&B, session->server_public_key, SRP_PRIME_BYTES);
  242. mbedtls_mpi_read_binary(&b, session->server_secret, SRP_PRIME_BYTES);
  243. // Check A != 0 and A % N != 0
  244. if (mbedtls_mpi_cmp_int(&A, 0) == 0) {
  245. ESP_LOGE(TAG, "Invalid client public key (zero)");
  246. goto cleanup;
  247. }
  248. mbedtls_mpi_mod_mpi(&tmp, &A, &N);
  249. if (mbedtls_mpi_cmp_int(&tmp, 0) == 0) {
  250. ESP_LOGE(TAG, "Invalid client public key (multiple of N)");
  251. goto cleanup;
  252. }
  253. // u = H(PAD(A) || PAD(B))
  254. {
  255. uint8_t ab_concat[SRP_PRIME_BYTES * 2];
  256. memcpy(ab_concat, session->client_public_key, SRP_PRIME_BYTES);
  257. memcpy(ab_concat + SRP_PRIME_BYTES, session->server_public_key,
  258. SRP_PRIME_BYTES);
  259. uint8_t u_hash[64];
  260. crypto_hash_sha512(u_hash, ab_concat, sizeof(ab_concat));
  261. mbedtls_mpi_read_binary(&u, u_hash, 64);
  262. }
  263. // Recompute k = H(N || pad(g))
  264. {
  265. uint8_t hash_input[SRP_PRIME_BYTES * 2];
  266. memcpy(hash_input, srp_N, SRP_PRIME_BYTES);
  267. memset(hash_input + SRP_PRIME_BYTES, 0, SRP_PRIME_BYTES);
  268. hash_input[SRP_PRIME_BYTES * 2 - 1] = SRP_GENERATOR;
  269. uint8_t k_hash[64];
  270. crypto_hash_sha512(k_hash, hash_input, sizeof(hash_input));
  271. mbedtls_mpi_read_binary(&k, k_hash, 64);
  272. mbedtls_mpi_mod_mpi(&k, &k, &N);
  273. }
  274. // Recompute x = H(s || H(I || ":" || P)) for "Pair-Setup:3939"
  275. {
  276. uint8_t inner_hash[64];
  277. crypto_hash_sha512_state state;
  278. crypto_hash_sha512_init(&state);
  279. crypto_hash_sha512_update(&state, (const uint8_t *)"Pair-Setup", 10);
  280. crypto_hash_sha512_update(&state, (const uint8_t *)":", 1);
  281. crypto_hash_sha512_update(&state, (const uint8_t *)"3939", 4);
  282. crypto_hash_sha512_final(&state, inner_hash);
  283. uint8_t x_hash[64];
  284. crypto_hash_sha512_init(&state);
  285. crypto_hash_sha512_update(&state, session->salt, SRP_SALT_BYTES);
  286. crypto_hash_sha512_update(&state, inner_hash, 64);
  287. crypto_hash_sha512_final(&state, x_hash);
  288. mbedtls_mpi_read_binary(&x, x_hash, 64);
  289. }
  290. // v = g^x mod N
  291. if (mbedtls_mpi_exp_mod(&v, &g, &x, &N, NULL) != 0) {
  292. goto cleanup;
  293. }
  294. // S = (A * v^u)^b mod N
  295. if (mbedtls_mpi_exp_mod(&tmp, &v, &u, &N, NULL) != 0) {
  296. goto cleanup;
  297. }
  298. if (mbedtls_mpi_mul_mpi(&tmp2, &A, &tmp) != 0) {
  299. goto cleanup;
  300. }
  301. mbedtls_mpi_mod_mpi(&tmp2, &tmp2, &N);
  302. if (mbedtls_mpi_exp_mod(&S, &tmp2, &b, &N, NULL) != 0) {
  303. goto cleanup;
  304. }
  305. // K = H(S)
  306. uint8_t S_bytes[SRP_PRIME_BYTES];
  307. size_t S_len = mpi_to_bytes_min(&S, S_bytes, sizeof(S_bytes));
  308. crypto_hash_sha512(session->session_key, S_bytes, S_len);
  309. session->session_key_len = 64;
  310. // Compute expected M1 = H(H(N)^H(g) || H(I) || s || A || B || K)
  311. uint8_t expected_m1[64];
  312. {
  313. // H(N)
  314. uint8_t h_N[64];
  315. crypto_hash_sha512(h_N, srp_N, sizeof(srp_N));
  316. // H(g)
  317. uint8_t g_byte = SRP_GENERATOR;
  318. uint8_t h_g[64];
  319. crypto_hash_sha512(h_g, &g_byte, 1);
  320. // H(N) ^ H(g)
  321. uint8_t h_Ng_xor[64];
  322. for (int i = 0; i < 64; i++) {
  323. h_Ng_xor[i] = h_N[i] ^ h_g[i];
  324. }
  325. // H(I) where I = "Pair-Setup"
  326. uint8_t h_I[64];
  327. crypto_hash_sha512(h_I, (const uint8_t *)"Pair-Setup", 10);
  328. // Get minimal representations
  329. const uint8_t *salt_ptr;
  330. size_t salt_len;
  331. trim_leading_zeros(session->salt, SRP_SALT_BYTES, &salt_ptr, &salt_len);
  332. uint8_t A_bytes[SRP_PRIME_BYTES];
  333. uint8_t B_bytes[SRP_PRIME_BYTES];
  334. size_t A_len = mpi_to_bytes_min(&A, A_bytes, sizeof(A_bytes));
  335. size_t B_len = mpi_to_bytes_min(&B, B_bytes, sizeof(B_bytes));
  336. compute_m1(expected_m1, h_Ng_xor, h_I, salt_ptr, salt_len, A_bytes, A_len,
  337. B_bytes, B_len, session->session_key, 64);
  338. }
  339. // Verify client proof
  340. if (memcmp(client_proof, expected_m1, SRP_PROOF_BYTES) != 0) {
  341. ESP_LOGE(TAG, "Client proof verification failed");
  342. goto cleanup;
  343. }
  344. memcpy(session->proof_m1, client_proof, SRP_PROOF_BYTES);
  345. {
  346. uint8_t A_bytes[SRP_PRIME_BYTES];
  347. size_t A_len = mpi_to_bytes_min(&A, A_bytes, sizeof(A_bytes));
  348. crypto_hash_sha512_state state;
  349. crypto_hash_sha512_init(&state);
  350. crypto_hash_sha512_update(&state, A_bytes, A_len);
  351. crypto_hash_sha512_update(&state, session->proof_m1, SRP_PROOF_BYTES);
  352. crypto_hash_sha512_update(&state, session->session_key,
  353. session->session_key_len);
  354. crypto_hash_sha512_final(&state, session->proof_m2);
  355. }
  356. session->verified = true;
  357. session->state = 2;
  358. ret = 0;
  359. cleanup:
  360. mbedtls_mpi_free(&N);
  361. mbedtls_mpi_free(&g);
  362. mbedtls_mpi_free(&A);
  363. mbedtls_mpi_free(&B);
  364. mbedtls_mpi_free(&b);
  365. mbedtls_mpi_free(&u);
  366. mbedtls_mpi_free(&S);
  367. mbedtls_mpi_free(&k);
  368. mbedtls_mpi_free(&v);
  369. mbedtls_mpi_free(&x);
  370. mbedtls_mpi_free(&tmp);
  371. mbedtls_mpi_free(&tmp2);
  372. return ret == 0 ? ESP_OK : ESP_FAIL;
  373. }
  374. const uint8_t *srp_get_proof(srp_session_t *session) {
  375. if (!session || !session->verified) {
  376. return NULL;
  377. }
  378. return session->proof_m2;
  379. }
  380. const uint8_t *srp_get_session_key(srp_session_t *session, size_t *len) {
  381. if (!session || !session->verified) {
  382. return NULL;
  383. }
  384. if (len) {
  385. *len = session->session_key_len;
  386. }
  387. return session->session_key;
  388. }