/* SPDX-License-Identifier: GPL-3.0-only */ #include #include #include #define MBEDTLS_SSL_EXTENDED_MASTER_SECRET #define MBEDTLS_SSL_EXTENDED_MS_ENABLED 1 #define MBEDTLS_ERR_ERROR_CORRUPTION_DETECTED -1 #define MBEDTLS_SSL_DEBUG_MSG(...) ((void)0) #define MBEDTLS_SSL_DEBUG_RET(...) ((void)0) #define MBEDTLS_SSL_DEBUG_BUF(...) ((void)0) typedef struct { int unused; } mbedtls_ssl_context; typedef struct { int resume, extended_ms; unsigned char randbytes[64], premaster[128]; size_t pmslen; int (*calc_verify)(const mbedtls_ssl_context *, unsigned char *, size_t *); int (*tls_prf)(const unsigned char *, size_t, const char *, const unsigned char *, size_t, unsigned char *, size_t); } mbedtls_ssl_handshake_params; static int hash_error, prf_error, hash_calls, prf_calls; static size_t hash_size, expected_seed; static void mbedtls_platform_zeroize(void *p, size_t n) { memset(p, 0, n); } static int verify(const mbedtls_ssl_context *ssl, unsigned char *out, size_t *len) { hash_calls++; if (hash_error) return hash_error; /* Deliberately leave seed_len=64. */ *len = hash_size; memset(out, 0x23, *len); return 0; } static int prf(const unsigned char *p, size_t n, const char *label, const unsigned char *seed, size_t len, unsigned char *out, size_t size) { prf_calls++; assert(len == expected_seed && size == 48); assert(strcmp(label, len == 64 ? "master secret" : "extended master secret") == 0); for (size_t i = 0; i < len; i++) assert(seed[i] == (len == 64 ? 0x45 : 0x23)); if (prf_error) return prf_error; memset(out, 0x67, size); return 0; } /* SDK_FUNCTIONS */ int main(void) { mbedtls_ssl_context ssl = {0}; for (unsigned sha = 0; sha < 2; sha++) { hash_size = sha ? 48 : 32; for (unsigned mode = 0; mode < 5; mode++) { mbedtls_ssl_handshake_params h = {.calc_verify = verify, .tls_prf = prf, .pmslen = 32, .extended_ms = mode != 3, .resume = mode == 4}; memset(h.premaster, 0xab, sizeof(h.premaster)); memset(h.randbytes, 0x45, sizeof(h.randbytes)); unsigned char master[48]; memset(master, 0xcd, sizeof(master)); hash_calls = prf_calls = 0; hash_error = mode == 0 ? -0x1234 : 0; prf_error = mode == 2 ? -0x2345 : 0; expected_seed = mode == 3 ? 64 : hash_size; int ret = ssl_compute_master(&h, master, &ssl); assert(ret == (mode == 0 ? hash_error : mode == 2 ? prf_error : 0)); assert(hash_calls == (mode == 3 || mode == 4 ? 0 : 1)); assert(prf_calls == (mode == 0 || mode == 4 ? 0 : 1)); for (size_t i = 0; i < sizeof(master); i++) assert(master[i] == (mode == 1 || mode == 3 ? 0x67 : 0xcd)); for (size_t i = 0; i < sizeof(h.premaster); i++) assert(h.premaster[i] == (mode == 1 || mode == 3 ? 0 : 0xab)); } } puts("EMS extracted master calculation: SHA256/SHA384 error, success, PRF failure, non-EMS, resumption PASS"); }