1use core::marker::PhantomData;
8use core::ops::Add;
9
10use derive_where::derive_where;
11use digest::block_api::{CoreProxy, SmallBlockSizeUser};
12use digest::{Output, OutputSizeUser};
13use generic_array::typenum::{IsLess, Le, NonZero, Sum, U256};
14use generic_array::{ArrayLength, GenericArray};
15use rand::{CryptoRng, Rng};
16use subtle::{ConstantTimeEq, CtOption};
17
18use super::{
19 Deserialize, GenerateKe1Result, GenerateKe2Result, GenerateKe3Result, KeyExchange, Serialize,
20 SerializedContext, SerializedCredentialRequest, SerializedCredentialResponse,
21 SerializedIdentifiers,
22};
23use crate::ciphersuite::{CipherSuite, KeGroup};
24use crate::errors::ProtocolError;
25use crate::hash::{Hash, OutputSize, ProxyHash};
26use crate::key_exchange::group::Group;
27use crate::key_exchange::shared::{self, NonceLen};
28pub use crate::key_exchange::shared::{DiffieHellman, Ke1Message, Ke1State};
29use crate::keypair::{PrivateKey, PublicKey};
30use crate::opaque::Identifiers;
31use crate::serialization::{ConcatExt, SliceExt};
32
33pub struct TripleDh<G, H>(PhantomData<(G, H)>);
49
50#[cfg_attr(
52 feature = "serde",
53 derive(serde::Deserialize, serde::Serialize),
54 serde(bound = "")
55)]
56#[derive_where(Clone, Debug, Eq, Hash, Ord, PartialEq, PartialOrd, ZeroizeOnDrop)]
57pub struct Ke2State<H: OutputSizeUser> {
58 pub(super) session_key: Output<H>,
59 pub(super) expected_mac: Output<H>,
60}
61
62#[cfg_attr(
64 feature = "serde",
65 derive(serde::Deserialize, serde::Serialize),
66 serde(bound(
67 deserialize = "H: serde::Deserialize<'de>, PublicKey<G>: serde::Deserialize<'de>",
68 serialize = "H: serde::Serialize, PublicKey<G>: serde::Serialize",
69 ))
70)]
71#[derive_where(Clone, ZeroizeOnDrop)]
72#[derive_where(Debug, Eq, Hash, PartialEq; H, PublicKey<G>)]
73pub struct Ke2Builder<G: Group, H: Hash>
74where
75 H::Core: ProxyHash,
76 <<H as CoreProxy>::Core as SmallBlockSizeUser>::_BlockSize: IsLess<U256>,
77 Le<<<H as CoreProxy>::Core as SmallBlockSizeUser>::_BlockSize, U256>: NonZero,
78 OutputSize<H>: ArrayLength,
79{
80 server_nonce: GenericArray<u8, NonceLen>,
81 transcript_hasher: H,
82 #[derive_where(skip(Zeroize))]
83 client_e_pk: PublicKey<G>,
84 #[derive_where(skip(Zeroize))]
85 server_e_pk: PublicKey<G>,
86 shared_secret_1: GenericArray<u8, G::PkLen>,
87 shared_secret_3: GenericArray<u8, G::PkLen>,
88}
89
90#[cfg_attr(
92 feature = "serde",
93 derive(serde::Deserialize, serde::Serialize),
94 serde(bound(
95 deserialize = "G::Pk: serde::Deserialize<'de>",
96 serialize = "G::Pk: serde::Serialize"
97 ))
98)]
99#[derive_where(Clone, ZeroizeOnDrop)]
100#[derive_where(Debug, Eq, Hash, Ord, PartialEq, PartialOrd; G::Pk)]
101pub struct Ke2Message<G: Group, H: Hash>
102where
103 H::Core: ProxyHash,
104 <<H as CoreProxy>::Core as SmallBlockSizeUser>::_BlockSize: IsLess<U256>,
105 Le<<<H as CoreProxy>::Core as SmallBlockSizeUser>::_BlockSize, U256>: NonZero,
106 OutputSize<H>: ArrayLength,
107{
108 pub(super) server_nonce: GenericArray<u8, NonceLen>,
109 #[derive_where(skip(Zeroize))]
110 pub(super) server_e_pk: PublicKey<G>,
111 pub(super) mac: Output<H>,
112}
113
114#[cfg_attr(
116 feature = "serde",
117 derive(serde::Deserialize, serde::Serialize),
118 serde(bound = "")
119)]
120#[derive_where(Clone, Debug, Eq, Hash, Ord, PartialEq, PartialOrd, ZeroizeOnDrop)]
121pub struct Ke3Message<H: Hash>
122where
123 H::Core: ProxyHash,
124 <<H as CoreProxy>::Core as SmallBlockSizeUser>::_BlockSize: IsLess<U256>,
125 Le<<<H as CoreProxy>::Core as SmallBlockSizeUser>::_BlockSize, U256>: NonZero,
126 OutputSize<H>: ArrayLength,
127{
128 pub(super) mac: Output<H>,
129}
130
131impl<G: Group + 'static, H: Hash> KeyExchange for TripleDh<G, H>
137where
138 G::Sk: DiffieHellman<G>,
139 H::Core: ProxyHash,
140 <<H as CoreProxy>::Core as SmallBlockSizeUser>::_BlockSize: IsLess<U256>,
141 Le<<<H as CoreProxy>::Core as SmallBlockSizeUser>::_BlockSize, U256>: NonZero,
142 OutputSize<H>: ArrayLength,
143{
144 type Group = G;
145 type Hash = H;
146
147 type KE1State = Ke1State<G>;
148 type KE2State<CS: CipherSuite> = Ke2State<H>;
149 type KE1Message = Ke1Message<G>;
150 type KE2Builder<'a, CS: CipherSuite<KeyExchange = Self>> = Ke2Builder<G, H>;
151 type KE2BuilderData<'a, CS: 'static + CipherSuite> = &'a PublicKey<G>;
152 type KE2BuilderInput<CS: CipherSuite> = GenericArray<u8, G::PkLen>;
153 type KE2Message = Ke2Message<G, H>;
154 type KE3Message = Ke3Message<H>;
155
156 fn generate_ke1<R: Rng + CryptoRng>(
157 rng: &mut R,
158 ) -> Result<GenerateKe1Result<Self>, ProtocolError> {
159 shared::generate_ke1(rng)
160 }
161
162 fn ke2_builder<'a, CS: CipherSuite<KeyExchange = Self>, R: Rng + CryptoRng>(
163 rng: &mut R,
164 credential_request: SerializedCredentialRequest<CS>,
165 ke1_message: Self::KE1Message,
166 credential_response: SerializedCredentialResponse<CS>,
167 client_s_pk: PublicKey<G>,
168 identifiers: SerializedIdentifiers<'_, KeGroup<CS>>,
169 context: SerializedContext<'a>,
170 ) -> Result<Self::KE2Builder<'a, CS>, ProtocolError> {
171 let shared::Ke2BuilderCommon {
172 server_nonce,
173 transcript_hasher,
174 client_e_pk,
175 server_e_pk,
176 shared_secret_1,
177 shared_secret_3,
178 } = shared::ke2_builder_common::<G, H, CS, R>(
179 rng,
180 credential_request,
181 ke1_message,
182 credential_response,
183 client_s_pk,
184 identifiers,
185 context,
186 )?;
187
188 Ok(Ke2Builder {
189 server_nonce,
190 transcript_hasher,
191 client_e_pk,
192 server_e_pk,
193 shared_secret_1,
194 shared_secret_3,
195 })
196 }
197
198 fn ke2_builder_data<'a, CS: 'static + CipherSuite<KeyExchange = Self>>(
199 builder: &'a Self::KE2Builder<'_, CS>,
200 ) -> Self::KE2BuilderData<'a, CS> {
201 &builder.client_e_pk
202 }
203
204 fn generate_ke2_input<CS: CipherSuite<KeyExchange = Self>, R: CryptoRng + Rng>(
205 builder: &Self::KE2Builder<'_, CS>,
206 _: &mut R,
207 server_s_sk: &PrivateKey<G>,
208 ) -> Self::KE2BuilderInput<CS> {
209 server_s_sk.ke_diffie_hellman(&builder.client_e_pk)
210 }
211
212 fn build_ke2<CS: CipherSuite<KeyExchange = Self>>(
213 mut builder: Self::KE2Builder<'_, CS>,
214 shared_secret_2: Self::KE2BuilderInput<CS>,
215 ) -> Result<GenerateKe2Result<CS>, ProtocolError> {
216 let transcript_digest = builder.transcript_hasher.clone().finalize();
217 let derived_keys = shared::derive_keys::<H>(
218 [
219 builder.shared_secret_1.as_slice(),
220 &shared_secret_2,
221 &builder.shared_secret_3,
222 ]
223 .into_iter(),
224 &transcript_digest,
225 )?;
226
227 let (mac, expected_mac) = shared::compute_ke2_macs(
228 &mut builder.transcript_hasher,
229 &derived_keys,
230 &transcript_digest,
231 )?;
232
233 Ok(GenerateKe2Result {
234 state: Ke2State {
235 session_key: derived_keys.session_key,
236 expected_mac,
237 },
238 message: Ke2Message {
239 server_nonce: builder.server_nonce,
240 server_e_pk: builder.server_e_pk.clone(),
241 mac,
242 },
243 #[cfg(test)]
244 handshake_secret: derived_keys.handshake_secret,
245 #[cfg(test)]
246 km2: derived_keys.km2,
247 })
248 }
249
250 fn generate_ke3<CS: CipherSuite<KeyExchange = Self>, R: CryptoRng + Rng>(
251 _: &mut R,
252 credential_request: SerializedCredentialRequest<CS>,
253 ke1_message: Self::KE1Message,
254 credential_response: SerializedCredentialResponse<CS>,
255 ke1_state: &Self::KE1State,
256 ke2_message: Self::KE2Message,
257 server_s_pk: PublicKey<G>,
258 client_s_sk: PrivateKey<G>,
259 identifiers: SerializedIdentifiers<'_, KeGroup<CS>>,
260 context: SerializedContext<'_>,
261 ) -> Result<GenerateKe3Result<Self>, ProtocolError> {
262 let mut transcript_hasher = shared::transcript(
263 &context,
264 &identifiers,
265 &credential_request,
266 &ke1_message.to_iter(),
267 &credential_response,
268 ke2_message.server_nonce,
269 &ke2_message.server_e_pk.serialize(),
270 );
271
272 let shared_secret_1 = ke1_state
273 .client_e_sk
274 .ke_diffie_hellman(&ke2_message.server_e_pk);
275 let shared_secret_2 = ke1_state.client_e_sk.ke_diffie_hellman(&server_s_pk);
276 let shared_secret_3 = client_s_sk.ke_diffie_hellman(&ke2_message.server_e_pk);
277
278 let (derived_keys, client_mac) = shared::finalize_ke3_transcript(
279 &mut transcript_hasher,
280 [
281 shared_secret_1.as_slice(),
282 shared_secret_2.as_slice(),
283 shared_secret_3.as_slice(),
284 ]
285 .into_iter(),
286 &ke2_message.mac,
287 )?;
288
289 Ok(GenerateKe3Result {
290 session_key: derived_keys.session_key,
291 message: Ke3Message { mac: client_mac },
292 #[cfg(test)]
293 handshake_secret: derived_keys.handshake_secret,
294 #[cfg(test)]
295 km3: derived_keys.km3,
296 })
297 }
298
299 fn finish_ke<CS: CipherSuite>(
300 ke2_state: &Self::KE2State<CS>,
301 ke3_message: Self::KE3Message,
302 _: Identifiers<'_>,
303 _: SerializedContext<'_>,
304 ) -> Result<Output<H>, ProtocolError> {
305 CtOption::new(
306 ke2_state.session_key.clone(),
307 ke2_state.expected_mac.ct_eq(&ke3_message.mac),
308 )
309 .into_option()
310 .ok_or(ProtocolError::InvalidLoginError)
311 }
312}
313
314impl<H: Hash> Deserialize for Ke2State<H>
320where
321 H::Core: ProxyHash,
322 <<H as CoreProxy>::Core as SmallBlockSizeUser>::_BlockSize: IsLess<U256>,
323 Le<<<H as CoreProxy>::Core as SmallBlockSizeUser>::_BlockSize, U256>: NonZero,
324 OutputSize<H>: ArrayLength,
325{
326 fn deserialize_take(input: &mut &[u8]) -> Result<Self, ProtocolError> {
327 Ok(Self {
328 session_key: input.take_array("session key")?.into_ha0_4(),
329 expected_mac: input.take_array("expected mac")?.into_ha0_4(),
330 })
331 }
332}
333
334impl<H: Hash> Serialize for Ke2State<H>
335where
336 H::Core: ProxyHash,
337 <<H as CoreProxy>::Core as SmallBlockSizeUser>::_BlockSize: IsLess<U256>,
338 Le<<<H as CoreProxy>::Core as SmallBlockSizeUser>::_BlockSize, U256>: NonZero,
339 OutputSize<H>: ArrayLength,
340 OutputSize<H>: Add<OutputSize<H>>,
342 Sum<OutputSize<H>, OutputSize<H>>: ArrayLength,
343{
344 type Len = Sum<OutputSize<H>, OutputSize<H>>;
345
346 fn serialize(&self) -> GenericArray<u8, Self::Len> {
347 let sk: GenericArray<u8, OutputSize<H>> =
348 GenericArray::from_slice(self.session_key.as_slice()).clone();
349 let mac: GenericArray<u8, OutputSize<H>> =
350 GenericArray::from_slice(self.expected_mac.as_slice()).clone();
351
352 sk.cat(mac)
353 }
354}
355
356impl<G: Group, H: Hash> Deserialize for Ke2Message<G, H>
357where
358 H::Core: ProxyHash,
359 <<H as CoreProxy>::Core as SmallBlockSizeUser>::_BlockSize: IsLess<U256>,
360 Le<<<H as CoreProxy>::Core as SmallBlockSizeUser>::_BlockSize, U256>: NonZero,
361 OutputSize<H>: ArrayLength,
362{
363 fn deserialize_take(input: &mut &[u8]) -> Result<Self, ProtocolError> {
364 Ok(Self {
365 server_nonce: input.take_array("server nonce")?,
366 server_e_pk: PublicKey::deserialize_take(input)?,
367 mac: input.take_array("mac")?.into_ha0_4(),
368 })
369 }
370}
371
372impl<H: Hash, G: Group> Serialize for Ke2Message<G, H>
373where
374 H::Core: ProxyHash,
375 <<H as CoreProxy>::Core as SmallBlockSizeUser>::_BlockSize: IsLess<U256>,
376 Le<<<H as CoreProxy>::Core as SmallBlockSizeUser>::_BlockSize, U256>: NonZero,
377 OutputSize<H>: ArrayLength,
378 NonceLen: Add<G::PkLen>,
380 Sum<NonceLen, G::PkLen>: ArrayLength + Add<OutputSize<H>>,
381 Sum<Sum<NonceLen, G::PkLen>, OutputSize<H>>: ArrayLength,
382{
383 type Len = Sum<Sum<NonceLen, G::PkLen>, OutputSize<H>>;
384
385 fn serialize(&self) -> GenericArray<u8, Self::Len> {
386 self.server_nonce
387 .cat(self.server_e_pk.serialize())
388 .cat(GenericArray::from_slice(self.mac.as_slice()).clone())
389 }
390}
391
392impl<H: Hash> Deserialize for Ke3Message<H>
393where
394 H::Core: ProxyHash,
395 <<H as CoreProxy>::Core as SmallBlockSizeUser>::_BlockSize: IsLess<U256>,
396 Le<<<H as CoreProxy>::Core as SmallBlockSizeUser>::_BlockSize, U256>: NonZero,
397 OutputSize<H>: ArrayLength,
398{
399 fn deserialize_take(bytes: &mut &[u8]) -> Result<Self, ProtocolError> {
400 Ok(Self {
401 mac: bytes.take_array("mac")?.into_ha0_4(),
402 })
403 }
404}
405
406impl<H: Hash> Serialize for Ke3Message<H>
407where
408 H::Core: ProxyHash,
409 <<H as CoreProxy>::Core as SmallBlockSizeUser>::_BlockSize: IsLess<U256>,
410 Le<<<H as CoreProxy>::Core as SmallBlockSizeUser>::_BlockSize, U256>: NonZero,
411 OutputSize<H>: ArrayLength,
412{
413 type Len = OutputSize<H>;
414
415 fn serialize(&self) -> GenericArray<u8, Self::Len> {
416 GenericArray::from_slice(self.mac.as_slice()).clone()
417 }
418}