Skip to main content

tls2/
crypto_default_provider.rs

1use crypto::{
2    Aead as CryptoAead, Hasher as CryptoHasher,
3    aes::{Aes128Gcm, Aes256Gcm},
4    chacha::ChaCha20Poly1305,
5    curve25519::x25519,
6    hkdf,
7    hmac::Hmac,
8    mlkem::{CIPHERTEXT_SIZE_768, generate_keypair_768_derand},
9    p256, random_fill,
10    sha2::{Sha256, Sha384},
11};
12use heapless::Vec;
13
14use crate::{
15    CipherSuite, CryptoProvider, Hash, KEY_EXCHANGE_PUBLIC_KEY_MAX_SIZE, KEY_EXCHANGE_SECRET_KEY_MAX_SIZE,
16    KEY_EXCHANGE_SHARED_SECRET_MAX_SIZE, KeyExchangeGroup, KeyExchangePublicKey, KeyExchangeSecretKey,
17    SIGNATURE_MAX_SIZE, SignatureScheme, errors::Error,
18};
19
20/// Default crypto provider backed by the `crypto` crate.
21///
22/// Provides AEAD, key exchange, hash, and signature operations.
23///
24/// # Examples
25///
26/// ```ignore
27/// use tls2::crypto_default_provider::DefaultCryptoProvider;
28///
29/// let crypto = DefaultCryptoProvider;
30/// ```
31#[derive(Clone)]
32pub struct DefaultCryptoProvider;
33
34/// Incremental hash state for the default provider.
35///
36/// Wraps either SHA-256 (32-bit state) or SHA-384 (64-bit state) depending
37/// on the negotiated cipher suite.
38#[derive(Clone)]
39pub enum Hasher {
40    Sha256(crypto::sha2::Sha256),
41    Sha384(crypto::sha2::Sha384),
42}
43
44/// Cached AEAD key for the default provider.
45///
46/// Wraps the expanded cipher state for AES-128-GCM, AES-256-GCM, or
47/// ChaCha20-Poly1305 so key expansion happens only once.
48#[cfg_attr(feature = "zeroize", derive(zeroize::Zeroize, zeroize::ZeroizeOnDrop))]
49pub enum AeadKey {
50    Aes128Gcm(Aes128Gcm),
51    Aes256Gcm(Aes256Gcm),
52    ChaCha20Poly1305(ChaCha20Poly1305),
53}
54
55impl CryptoProvider for DefaultCryptoProvider {
56    type Hasher = Hasher;
57    type AeadKey = AeadKey;
58
59    #[inline]
60    fn cipher_suites() -> &'static [CipherSuite] {
61        // runtime / compile detection of CPU instructions to return the list of supported ciphersuites
62        // in the correct order
63        static AES_FIRST: &[CipherSuite] = &[
64            CipherSuite::TlsAes256GcmSha384,
65            CipherSuite::TlsChaCha20Poly1305Sha256,
66            CipherSuite::TlsAes128GcmSha256,
67        ];
68        static CHACHA_FIRST: &[CipherSuite] = &[
69            CipherSuite::TlsChaCha20Poly1305Sha256,
70            CipherSuite::TlsAes256GcmSha384,
71            CipherSuite::TlsAes128GcmSha256,
72        ];
73
74        #[cfg(target_arch = "x86_64")]
75        {
76            #[cfg(feature = "std")]
77            if std::arch::is_x86_feature_detected!("aes") {
78                return AES_FIRST;
79            }
80            #[cfg(all(not(feature = "std"), target_feature = "aes"))]
81            return AES_FIRST;
82        }
83
84        #[cfg(target_arch = "aarch64")]
85        {
86            #[cfg(feature = "std")]
87            if std::arch::is_aarch64_feature_detected!("aes") {
88                return AES_FIRST;
89            }
90            #[cfg(all(not(feature = "std"), target_feature = "aes"))]
91            return AES_FIRST;
92        }
93
94        CHACHA_FIRST
95    }
96
97    #[inline]
98    fn signature_schemes() -> &'static [SignatureScheme] {
99        &[
100            SignatureScheme::Ed25519,
101            SignatureScheme::EcdsaP256Sha256,
102            SignatureScheme::EcdsaP384Sha384,
103            SignatureScheme::RsaPssRsaSha256,
104            SignatureScheme::RsaPkcs1Sha256,
105        ]
106    }
107
108    #[inline]
109    fn key_exchange_groups() -> &'static [KeyExchangeGroup] {
110        &[
111            KeyExchangeGroup::X25519MlKem768,
112            KeyExchangeGroup::X25519,
113            KeyExchangeGroup::Secp256r1,
114        ]
115    }
116
117    fn new_aead_key(&self, suite: CipherSuite, key: &[u8]) -> Self::AeadKey {
118        match suite {
119            CipherSuite::TlsAes128GcmSha256 => AeadKey::Aes128Gcm(Aes128Gcm::new(key.try_into().unwrap())),
120            CipherSuite::TlsAes256GcmSha384 => AeadKey::Aes256Gcm(Aes256Gcm::new(key.try_into().unwrap())),
121            CipherSuite::TlsChaCha20Poly1305Sha256 => {
122                AeadKey::ChaCha20Poly1305(ChaCha20Poly1305::new(key.try_into().unwrap()))
123            }
124        }
125    }
126
127    fn new_hash(&self, suite: CipherSuite) -> Self::Hasher {
128        match suite {
129            CipherSuite::TlsAes128GcmSha256 | CipherSuite::TlsChaCha20Poly1305Sha256 => {
130                Hasher::Sha256(<Sha256 as CryptoHasher>::new())
131            }
132            CipherSuite::TlsAes256GcmSha384 => Hasher::Sha384(<Sha384 as CryptoHasher>::new()),
133        }
134    }
135
136    fn hash_update(&self, state: &mut Self::Hasher, data: &[u8]) {
137        match state {
138            Hasher::Sha256(s) => s.update(data),
139            Hasher::Sha384(s) => s.update(data),
140        }
141    }
142
143    fn hash_finalize(&self, state: Self::Hasher) -> Result<Hash, Error> {
144        match state {
145            Hasher::Sha256(s) => {
146                let h = s.sum();
147                Ok(Hash::from_slice(h.as_ref()))
148            }
149            Hasher::Sha384(s) => {
150                let h = s.sum();
151                Ok(Hash::from_slice(h.as_ref()))
152            }
153        }
154    }
155
156    fn secure_random(&self, buf: &mut [u8]) {
157        random_fill(buf);
158    }
159
160    fn hash(&self, suite: CipherSuite, data: &[u8]) -> Result<Hash, Error> {
161        match suite {
162            CipherSuite::TlsAes128GcmSha256 | CipherSuite::TlsChaCha20Poly1305Sha256 => {
163                Ok(Hash::from_slice(Sha256::hash(data).as_ref()))
164            }
165            CipherSuite::TlsAes256GcmSha384 => Ok(Hash::from_slice(Sha384::hash(data).as_ref())),
166        }
167    }
168
169    fn hmac(&self, suite: CipherSuite, key: &Hash, data: &[u8]) -> Result<Hash, Error> {
170        match suite {
171            CipherSuite::TlsAes128GcmSha256 | CipherSuite::TlsChaCha20Poly1305Sha256 => {
172                Ok(Hash::from_slice(Hmac::<Sha256>::mac(key, data).as_ref()))
173            }
174            CipherSuite::TlsAes256GcmSha384 => Ok(Hash::from_slice(Hmac::<Sha384>::mac(key, data).as_ref())),
175        }
176    }
177
178    fn hkdf_extract(&self, suite: CipherSuite, salt: &Hash, ikm: &[u8]) -> Result<Hash, Error> {
179        match suite {
180            CipherSuite::TlsAes128GcmSha256 | CipherSuite::TlsChaCha20Poly1305Sha256 => {
181                Ok(Hash::from_slice(hkdf::extract::<Sha256>(Some(salt), ikm).as_ref()))
182            }
183            CipherSuite::TlsAes256GcmSha384 => Ok(Hash::from_slice(hkdf::extract::<Sha384>(Some(salt), ikm).as_ref())),
184        }
185    }
186
187    fn hkdf_expand_label(
188        &self,
189        out: &mut [u8],
190        suite: CipherSuite,
191        secret: &Hash,
192        label: &[u8],
193        context: &[u8],
194    ) -> Result<(), Error> {
195        let len = out.len();
196        let hkdf_label = build_hkdf_label(label, context, len);
197        match suite {
198            CipherSuite::TlsAes128GcmSha256 | CipherSuite::TlsChaCha20Poly1305Sha256 => {
199                hkdf::expand::<Sha256>(out, secret, &hkdf_label).unwrap();
200            }
201            CipherSuite::TlsAes256GcmSha384 => {
202                hkdf::expand::<Sha384>(out, secret, &hkdf_label).unwrap();
203            }
204        }
205        Ok(())
206    }
207
208    fn aead_encrypt(
209        &self,
210        key: &Self::AeadKey,
211        nonce: &[u8],
212        aad: &[u8],
213        data: &mut [u8],
214        plaintext_len: usize,
215    ) -> Result<usize, Error> {
216        let total = plaintext_len + 16;
217        match key {
218            AeadKey::Aes128Gcm(cipher) => {
219                let tag = cipher.encrypt_in_place(&mut data[..plaintext_len], nonce, aad);
220                data[plaintext_len..total].copy_from_slice(tag.as_ref());
221            }
222            AeadKey::Aes256Gcm(cipher) => {
223                let tag = cipher.encrypt_in_place(&mut data[..plaintext_len], nonce, aad);
224                data[plaintext_len..total].copy_from_slice(tag.as_ref());
225            }
226            AeadKey::ChaCha20Poly1305(cipher) => {
227                let tag = cipher.encrypt_in_place(&mut data[..plaintext_len], nonce, aad);
228                data[plaintext_len..total].copy_from_slice(tag.as_ref());
229            }
230        }
231        Ok(total)
232    }
233
234    fn aead_decrypt(&self, key: &Self::AeadKey, nonce: &[u8], aad: &[u8], data: &mut [u8]) -> Result<usize, Error> {
235        if data.len() < 16 {
236            return Err(Error::AeadError);
237        }
238        let ct_len = data.len() - 16;
239        let (ct, tag) = data.split_at_mut(ct_len);
240        match key {
241            AeadKey::Aes128Gcm(cipher) => {
242                cipher
243                    .decrypt_in_place(ct, nonce, aad, tag)
244                    .map_err(|_| Error::AeadError)?;
245            }
246            AeadKey::Aes256Gcm(cipher) => {
247                cipher
248                    .decrypt_in_place(ct, nonce, aad, tag)
249                    .map_err(|_| Error::AeadError)?;
250            }
251            AeadKey::ChaCha20Poly1305(cipher) => {
252                cipher
253                    .decrypt_in_place(ct, nonce, aad, tag)
254                    .map_err(|_| Error::AeadError)?;
255            }
256        }
257        Ok(ct_len)
258    }
259
260    fn key_exchange_generate_keypair(
261        &self,
262        group: KeyExchangeGroup,
263    ) -> Result<(KeyExchangeSecretKey, KeyExchangePublicKey), Error> {
264        match group {
265            KeyExchangeGroup::X25519 => {
266                let sk = x25519::SecretKey::generate();
267                let pk = sk.public_key();
268                Ok((
269                    KeyExchangeSecretKey::new(group, &sk.to_bytes()),
270                    KeyExchangePublicKey::new(group, &pk.to_bytes()),
271                ))
272            }
273            KeyExchangeGroup::X25519MlKem768 => {
274                // for X25519MlKem768 we keep the same order as for the public keys and shared secret:
275                // first the ml_kem seed (64 bytes) and then the x25519 seed (32 bytes)
276                let mut seeds = [0u8; KEY_EXCHANGE_SECRET_KEY_MAX_SIZE];
277                self.secure_random(&mut seeds);
278
279                let (_, mlkem_pk) =
280                    generate_keypair_768_derand(&seeds[..64].try_into().map_err(|_| Error::CryptoError)?);
281                let x25519_sk =
282                    x25519::SecretKey::from_bytes(&seeds[64..96].try_into().map_err(|_| Error::CryptoError)?);
283                let x25519_pk = x25519_sk.public_key();
284
285                let mut pub_bytes: Vec<u8, KEY_EXCHANGE_PUBLIC_KEY_MAX_SIZE> = Vec::new();
286                pub_bytes
287                    .extend_from_slice(&mlkem_pk.to_bytes())
288                    .map_err(|_| Error::CryptoError)?;
289                pub_bytes
290                    .extend_from_slice(&x25519_pk.to_bytes())
291                    .map_err(|_| Error::CryptoError)?;
292                Ok((
293                    KeyExchangeSecretKey::new(group, &seeds),
294                    KeyExchangePublicKey::new(group, &pub_bytes),
295                ))
296            }
297            KeyExchangeGroup::Secp256r1 => {
298                let mut private_key_bytes = [0u8; 32];
299                self.secure_random(&mut private_key_bytes);
300                let private_key = p256::SecretKey::from_bytes(&private_key_bytes).map_err(|_| Error::CryptoError)?;
301                let public_key = private_key.public_key();
302                Ok((
303                    KeyExchangeSecretKey::new(group, &private_key.to_bytes()),
304                    KeyExchangePublicKey::new(group, &public_key.to_bytes()),
305                ))
306            }
307            _ => Err(Error::UnsupportedKeyExchangeGroup),
308        }
309    }
310
311    fn key_exchange(
312        &self,
313        secret: &KeyExchangeSecretKey,
314        peer_public: &[u8],
315    ) -> Result<Vec<u8, KEY_EXCHANGE_SHARED_SECRET_MAX_SIZE>, Error> {
316        match secret.group() {
317            KeyExchangeGroup::X25519 => {
318                let sk_bytes: &[u8; 32] = secret.bytes().try_into().map_err(|_| Error::CryptoError)?;
319                let pk_bytes: &[u8; 32] = peer_public.try_into().map_err(|_| Error::CryptoError)?;
320                let sk = x25519::SecretKey::from_bytes(sk_bytes);
321                let pk = x25519::PublicKey::from_bytes(pk_bytes);
322                let ss = sk.ecdh(&pk);
323                let mut out = Vec::new();
324                out.extend_from_slice(&ss).unwrap();
325                Ok(out)
326            }
327            KeyExchangeGroup::X25519MlKem768 => {
328                let seeds: &[u8; KEY_EXCHANGE_SECRET_KEY_MAX_SIZE] =
329                    secret.bytes().try_into().map_err(|_| Error::CryptoError)?;
330
331                let (mlkem_sk, _) =
332                    generate_keypair_768_derand(&seeds[..64].try_into().map_err(|_| Error::CryptoError)?);
333                let x25519_sk =
334                    x25519::SecretKey::from_bytes(&seeds[64..96].try_into().map_err(|_| Error::CryptoError)?);
335
336                let mlkem_ct: &[u8; CIPHERTEXT_SIZE_768] = peer_public
337                    .get(..CIPHERTEXT_SIZE_768)
338                    .and_then(|s| s.try_into().ok())
339                    .ok_or(Error::CryptoError)?;
340                let peer_x25519_pk = x25519::PublicKey::from_bytes(
341                    &peer_public[CIPHERTEXT_SIZE_768..]
342                        .try_into()
343                        .map_err(|_| Error::CryptoError)?,
344                );
345                let shared_secret_mlkem = mlkem_sk.decapsulate(mlkem_ct).map_err(|_| Error::CryptoError)?;
346                let shared_secret_x25519 = x25519_sk.ecdh(&peer_x25519_pk);
347                let mut out = Vec::new();
348                out.extend_from_slice(&shared_secret_mlkem)
349                    .map_err(|_| Error::CryptoError)?;
350                out.extend_from_slice(&shared_secret_x25519)
351                    .map_err(|_| Error::CryptoError)?;
352                Ok(out)
353            }
354            KeyExchangeGroup::Secp256r1 => {
355                let sk_bytes: &[u8; 32] = secret.bytes().try_into().map_err(|_| Error::CryptoError)?;
356                let private_key = p256::SecretKey::from_bytes(sk_bytes).map_err(|_| Error::CryptoError)?;
357                let shared_secret = private_key
358                    .ecdh(&p256::PublicKey::from_bytes(peer_public).map_err(|_| Error::CryptoError)?)
359                    .map_err(|_| Error::CryptoError)?;
360                let mut out = Vec::new();
361                out.extend_from_slice(&shared_secret).unwrap();
362                Ok(out)
363            }
364            _ => Err(Error::UnsupportedKeyExchangeGroup),
365        }
366    }
367
368    fn sign(
369        &self,
370        scheme: SignatureScheme,
371        secret_key: &[u8],
372        data: &[u8],
373    ) -> Result<Vec<u8, SIGNATURE_MAX_SIZE>, Error> {
374        match scheme {
375            SignatureScheme::Ed25519 => {
376                let seed: &[u8; 32] = secret_key.try_into().map_err(|_| Error::CryptoError)?;
377                let secret_key = crypto::curve25519::ed25519::SecretKey::from_bytes(seed);
378                let signature = secret_key.sign(data);
379                Ok(Vec::from_slice(&signature).unwrap())
380            }
381            SignatureScheme::EcdsaP256Sha256 => {
382                let key: &[u8; 32] = secret_key.try_into().map_err(|_| Error::CryptoError)?;
383                let private_key = crypto::p256::SecretKey::from_bytes(key).map_err(|_| Error::CryptoError)?;
384                let raw_sig = private_key.sign(data).map_err(|_| Error::CryptoError)?;
385                p256_raw_to_der_signature(&raw_sig)
386            }
387            SignatureScheme::EcdsaP384Sha384 => {
388                let key: &[u8; 48] = secret_key.try_into().map_err(|_| Error::CryptoError)?;
389                let private_key = crypto::p384::PrivateKey::from_bytes(key).map_err(|_| Error::CryptoError)?;
390                let raw_sig = private_key.sign(data).map_err(|_| Error::CryptoError)?;
391                p384_raw_to_der_signature(&raw_sig)
392            }
393            _ => Err(Error::CryptoError),
394        }
395    }
396
397    fn verify(&self, scheme: SignatureScheme, public_key: &[u8], data: &[u8], signature: &[u8]) -> Result<(), Error> {
398        match scheme {
399            SignatureScheme::Ed25519 => {
400                let pk: &[u8; 32] = public_key.try_into().map_err(|_| Error::CryptoError)?;
401                let pk = crypto::curve25519::ed25519::PublicKey::from_bytes(pk).map_err(|_| Error::CryptoError)?;
402                let sig: &[u8; 64] = signature.try_into().map_err(|_| Error::InvalidSignature)?;
403                pk.verify(data, sig).map_err(|_| Error::InvalidSignature)
404            }
405            SignatureScheme::EcdsaP256Sha256 => {
406                let pk = crypto::p256::PublicKey::from_bytes(public_key).map_err(|_| Error::CryptoError)?;
407                let sig = p256_der_to_raw_signature(signature).map_err(|_| Error::DecodeError)?;
408                pk.verify(data, &sig).map_err(|_| Error::InvalidSignature)
409            }
410            SignatureScheme::EcdsaP384Sha384 => {
411                let pk = crypto::p384::PublicKey::from_bytes(public_key).map_err(|_| Error::CryptoError)?;
412                let sig = p384_der_to_raw_signature(signature).map_err(|_| Error::DecodeError)?;
413                pk.verify(data, &sig).map_err(|_| Error::InvalidSignature)
414            }
415            SignatureScheme::RsaPkcs1Sha256 => {
416                crypto::rsa::verify_pkcs1_sha256(public_key, signature, data).map_err(|_| Error::InvalidSignature)
417            }
418            SignatureScheme::RsaPssRsaSha256 => {
419                crypto::rsa::verify_pss_sha256(public_key, signature, data).map_err(|_| Error::InvalidSignature)
420            }
421            _ => Err(Error::InvalidSignature),
422        }
423    }
424}
425
426// ── HKDF helpers ──
427
428fn build_hkdf_label(label: &[u8], context: &[u8], out_len: usize) -> Vec<u8, 130> {
429    let mut buf = Vec::new();
430    buf.extend_from_slice(&(out_len as u16).to_be_bytes()).unwrap();
431    buf.push((6 + label.len()) as u8).unwrap();
432    buf.extend_from_slice(b"tls13 ").unwrap();
433    buf.extend_from_slice(label).unwrap();
434    buf.push(context.len() as u8).unwrap();
435    buf.extend_from_slice(context).unwrap();
436    buf
437}
438
439fn normalize_scalar(bytes: &[u8]) -> [u8; 32] {
440    let mut out = [0u8; 32];
441    let start = bytes.len().saturating_sub(32);
442    let copy_len = bytes.len().saturating_sub(start);
443    out[32 - copy_len..].copy_from_slice(&bytes[start..]);
444    out
445}
446
447fn normalize_scalar_48(bytes: &[u8]) -> [u8; 48] {
448    let mut out = [0u8; 48];
449    let start = bytes.len().saturating_sub(48);
450    let copy_len = bytes.len().saturating_sub(start);
451    out[48 - copy_len..].copy_from_slice(&bytes[start..]);
452    out
453}
454
455fn p256_der_to_raw_signature(der: &[u8]) -> Result<[u8; 64], Error> {
456    if der.len() < 8 || der[0] != 0x30 {
457        return Err(Error::DecodeError);
458    }
459    let mut pos = 2;
460    if pos + 2 > der.len() || der[pos] != 0x02 {
461        return Err(Error::DecodeError);
462    }
463    pos += 1;
464    let r_len = der[pos] as usize;
465    pos += 1;
466    if pos + r_len > der.len() {
467        return Err(Error::DecodeError);
468    }
469    let r = normalize_scalar(&der[pos..pos + r_len]);
470    pos += r_len;
471    if pos + 2 > der.len() || der[pos] != 0x02 {
472        return Err(Error::DecodeError);
473    }
474    pos += 1;
475    let s_len = der[pos] as usize;
476    pos += 1;
477    if pos + s_len > der.len() {
478        return Err(Error::DecodeError);
479    }
480    let s = normalize_scalar(&der[pos..pos + s_len]);
481    let mut sig = [0u8; 64];
482    sig[..32].copy_from_slice(&r);
483    sig[32..].copy_from_slice(&s);
484    Ok(sig)
485}
486
487fn p384_der_to_raw_signature(der: &[u8]) -> Result<[u8; 96], Error> {
488    if der.len() < 8 || der[0] != 0x30 {
489        return Err(Error::DecodeError);
490    }
491    let mut pos = 2;
492    if pos + 2 > der.len() || der[pos] != 0x02 {
493        return Err(Error::DecodeError);
494    }
495    pos += 1;
496    let r_len = der[pos] as usize;
497    pos += 1;
498    if pos + r_len > der.len() {
499        return Err(Error::DecodeError);
500    }
501    let r = normalize_scalar_48(&der[pos..pos + r_len]);
502    pos += r_len;
503    if pos + 2 > der.len() || der[pos] != 0x02 {
504        return Err(Error::DecodeError);
505    }
506    pos += 1;
507    let s_len = der[pos] as usize;
508    pos += 1;
509    if pos + s_len > der.len() {
510        return Err(Error::DecodeError);
511    }
512    let s = normalize_scalar_48(&der[pos..pos + s_len]);
513    let mut sig = [0u8; 96];
514    sig[..48].copy_from_slice(&r);
515    sig[48..].copy_from_slice(&s);
516    Ok(sig)
517}
518
519fn p256_raw_to_der_signature(raw: &[u8; 64]) -> Result<Vec<u8, SIGNATURE_MAX_SIZE>, Error> {
520    let r = &raw[..32];
521    let s = &raw[32..];
522
523    let r_offset = r.iter().position(|&b| b != 0).unwrap_or(r.len());
524    let r_stripped = &r[r_offset..];
525    let (r_enc, r_enc_len) = if r_stripped.is_empty() {
526        ([0x00u8; 34], 1)
527    } else if r_stripped[0] & 0x80 != 0 {
528        let mut arr = [0u8; 34];
529        arr[0] = 0x00;
530        arr[1..=r_stripped.len()].copy_from_slice(&r_stripped);
531        (arr, r_stripped.len() + 1)
532    } else {
533        let mut arr = [0u8; 34];
534        arr[..r_stripped.len()].copy_from_slice(&r_stripped);
535        (arr, r_stripped.len())
536    };
537
538    let s_offset = s.iter().position(|&b| b != 0).unwrap_or(s.len());
539    let s_stripped = &s[s_offset..];
540    let (s_enc, s_enc_len) = if s_stripped.is_empty() {
541        ([0x00u8; 34], 1)
542    } else if s_stripped[0] & 0x80 != 0 {
543        let mut arr = [0u8; 34];
544        arr[0] = 0x00;
545        arr[1..=s_stripped.len()].copy_from_slice(&s_stripped);
546        (arr, s_stripped.len() + 1)
547    } else {
548        let mut arr = [0u8; 34];
549        arr[..s_stripped.len()].copy_from_slice(&s_stripped);
550        (arr, s_stripped.len())
551    };
552
553    let total_len = 2 + r_enc_len + 2 + s_enc_len;
554    let mut der = Vec::new();
555    der.push(0x30).unwrap();
556    der.push(total_len as u8).unwrap();
557    der.push(0x02).unwrap();
558    der.push(r_enc_len as u8).unwrap();
559    der.extend_from_slice(&r_enc[..r_enc_len]).unwrap();
560    der.push(0x02).unwrap();
561    der.push(s_enc_len as u8).unwrap();
562    der.extend_from_slice(&s_enc[..s_enc_len]).unwrap();
563    Ok(der)
564}
565
566fn p384_raw_to_der_signature(raw: &[u8; 96]) -> Result<Vec<u8, SIGNATURE_MAX_SIZE>, Error> {
567    let r = &raw[..48];
568    let s = &raw[48..];
569
570    let r_offset = r.iter().position(|&b| b != 0).unwrap_or(r.len());
571    let r_stripped = &r[r_offset..];
572    let (r_enc, r_enc_len) = if r_stripped.is_empty() {
573        ([0x00u8; 50], 1)
574    } else if r_stripped[0] & 0x80 != 0 {
575        let mut arr = [0u8; 50];
576        arr[0] = 0x00;
577        arr[1..=r_stripped.len()].copy_from_slice(r_stripped);
578        (arr, r_stripped.len() + 1)
579    } else {
580        let mut arr = [0u8; 50];
581        arr[..r_stripped.len()].copy_from_slice(r_stripped);
582        (arr, r_stripped.len())
583    };
584
585    let s_offset = s.iter().position(|&b| b != 0).unwrap_or(s.len());
586    let s_stripped = &s[s_offset..];
587    let (s_enc, s_enc_len) = if s_stripped.is_empty() {
588        ([0x00u8; 50], 1)
589    } else if s_stripped[0] & 0x80 != 0 {
590        let mut arr = [0u8; 50];
591        arr[0] = 0x00;
592        arr[1..=s_stripped.len()].copy_from_slice(s_stripped);
593        (arr, s_stripped.len() + 1)
594    } else {
595        let mut arr = [0u8; 50];
596        arr[..s_stripped.len()].copy_from_slice(s_stripped);
597        (arr, s_stripped.len())
598    };
599
600    let total_len = 2 + r_enc_len + 2 + s_enc_len;
601    let mut der = Vec::new();
602    der.push(0x30).unwrap();
603    der.push(total_len as u8).unwrap();
604    der.push(0x02).unwrap();
605    der.push(r_enc_len as u8).unwrap();
606    der.extend_from_slice(&r_enc[..r_enc_len]).unwrap();
607    der.push(0x02).unwrap();
608    der.push(s_enc_len as u8).unwrap();
609    der.extend_from_slice(&s_enc[..s_enc_len]).unwrap();
610    Ok(der)
611}
612
613#[cfg(test)]
614mod tests {
615    use crypto::{
616        curve25519::x25519,
617        mlkem::{PUBLIC_KEY_SIZE_768, PublicKey768},
618    };
619
620    use super::*;
621
622    #[test]
623    fn x25519_mlkem768_key_exchange() {
624        let provider = DefaultCryptoProvider;
625        let group = KeyExchangeGroup::X25519MlKem768;
626
627        // ── Client side: generate keypair, send ClientHello ──
628        let (client_secret, client_public) = provider.key_exchange_generate_keypair(group).unwrap();
629        assert_eq!(client_public.bytes().len(), 1216, "client key share must be 1216 bytes");
630
631        // ── Server side: parse client's share, encapsulate ──
632        let client_mlkem_pk =
633            PublicKey768::from_bytes(client_public.bytes()[..PUBLIC_KEY_SIZE_768].try_into().unwrap());
634        let (server_ct, server_mlkem_ss) = client_mlkem_pk.encapsulate();
635
636        let server_x25519_sk = x25519::SecretKey::generate();
637        let server_x25519_pk = server_x25519_sk.public_key();
638        let server_x25519_ss = server_x25519_sk.ecdh(&x25519::PublicKey::from_bytes(
639            &client_public.bytes()[PUBLIC_KEY_SIZE_768..].try_into().unwrap(),
640        ));
641        let server_ss_full: [u8; 64] = {
642            let mut s = [0u8; 64];
643            s[..32].copy_from_slice(&server_mlkem_ss);
644            s[32..].copy_from_slice(&server_x25519_ss);
645            s
646        };
647
648        // Build server's key share (1120 bytes: ct + X25519 pk)
649        let mut server_share: Vec<u8, KEY_EXCHANGE_PUBLIC_KEY_MAX_SIZE> = Vec::new();
650        server_share.extend_from_slice(&server_ct).unwrap();
651        server_share.extend_from_slice(&server_x25519_pk.to_bytes()).unwrap();
652        assert_eq!(server_share.len(), 1120, "server key share must be 1120 bytes");
653
654        // ── Client side: compute shared secret from server's response ──
655        let client_ss = provider.key_exchange(&client_secret, &server_share).unwrap();
656        assert_eq!(client_ss.len(), 64, "shared secret must be 64 bytes");
657        assert_eq!(&client_ss[..], &server_ss_full, "shared secrets must match");
658    }
659
660    #[test]
661    fn secp256r1_key_exchange() {
662        let provider = DefaultCryptoProvider;
663        let group = KeyExchangeGroup::Secp256r1;
664
665        // Two parties each generate a keypair
666        let (alice_secret, alice_public) = provider.key_exchange_generate_keypair(group).unwrap();
667        let (bob_secret, bob_public) = provider.key_exchange_generate_keypair(group).unwrap();
668
669        assert_eq!(alice_public.bytes().len(), 65, "P-256 public key must be 65 bytes");
670        assert_eq!(bob_public.bytes().len(), 65, "P-256 public key must be 65 bytes");
671
672        // Each computes the shared secret from the other's public key
673        let alice_ss = provider.key_exchange(&alice_secret, bob_public.bytes()).unwrap();
674        let bob_ss = provider.key_exchange(&bob_secret, alice_public.bytes()).unwrap();
675
676        assert_eq!(alice_ss.len(), 32, "P-256 shared secret must be 32 bytes");
677        assert_eq!(alice_ss, bob_ss, "shared secrets must match");
678    }
679
680    #[test]
681    fn ecdsa_p256_sign_verify_roundtrip() {
682        let provider = DefaultCryptoProvider;
683        let data = b"TLS 1.3 test message";
684
685        let seed = [42u8; 32];
686        let private_key = crypto::p256::SecretKey::from_bytes(&seed).unwrap();
687        let public_key = private_key.public_key().to_bytes();
688
689        let signature = provider.sign(SignatureScheme::EcdsaP256Sha256, &seed, data).unwrap();
690        provider
691            .verify(SignatureScheme::EcdsaP256Sha256, &public_key, data, &signature)
692            .unwrap();
693    }
694
695    #[test]
696    fn ecdsa_p384_sign_verify_roundtrip() {
697        let provider = DefaultCryptoProvider;
698        let data = b"TLS 1.3 test message";
699
700        let seed = [42u8; 48];
701        let private_key = crypto::p384::PrivateKey::from_bytes(&seed).unwrap();
702        let public_key = private_key.public_key().to_bytes();
703
704        let signature = provider.sign(SignatureScheme::EcdsaP384Sha384, &seed, data).unwrap();
705        provider
706            .verify(SignatureScheme::EcdsaP384Sha384, &public_key, data, &signature)
707            .unwrap();
708    }
709
710    #[test]
711    fn ecdsa_p256_der_encoding_roundtrip() {
712        let raw = [0u8; 64];
713        let der = p256_raw_to_der_signature(&raw).unwrap();
714        let decoded = p256_der_to_raw_signature(&der).unwrap();
715        assert_eq!(raw, decoded);
716
717        let raw = [0xffu8; 64];
718        let der = p256_raw_to_der_signature(&raw).unwrap();
719        let decoded = p256_der_to_raw_signature(&der).unwrap();
720        assert_eq!(raw, decoded);
721
722        let mut raw = [0u8; 64];
723        raw[0] = 0x12;
724        raw[32] = 0x34;
725        let der = p256_raw_to_der_signature(&raw).unwrap();
726        let decoded = p256_der_to_raw_signature(&der).unwrap();
727        assert_eq!(raw, decoded);
728    }
729
730    #[test]
731    fn ecdsa_p384_der_encoding_roundtrip() {
732        let raw = [0u8; 96];
733        let der = p384_raw_to_der_signature(&raw).unwrap();
734        let decoded = p384_der_to_raw_signature(&der).unwrap();
735        assert_eq!(raw, decoded);
736
737        let raw = [0xffu8; 96];
738        let der = p384_raw_to_der_signature(&raw).unwrap();
739        let decoded = p384_der_to_raw_signature(&der).unwrap();
740        assert_eq!(raw, decoded);
741
742        let mut raw = [0u8; 96];
743        raw[0] = 0x12;
744        raw[48] = 0x34;
745        let der = p384_raw_to_der_signature(&raw).unwrap();
746        let decoded = p384_der_to_raw_signature(&der).unwrap();
747        assert_eq!(raw, decoded);
748    }
749}