1use core::ops::Add;
6
7use derive_where::derive_where;
8use digest::{FixedOutput, Output, Update};
9use generic_array::typenum::Sum;
10use generic_array::{ArrayLength, GenericArray};
11use zeroize::Zeroize;
12
13use crate::ciphersuite::{CipherSuite, KeGroup, KeHash, OprfGroup};
14use crate::errors::ProtocolError;
15use crate::hash::OutputSize;
16use crate::key_exchange::group::Group;
17use crate::key_exchange::shared::{Ke1MessageIter, Ke1MessageIterLen, NonceLen};
18use crate::key_exchange::{
19 Deserialize, Serialize, SerializedContext, SerializedCredentialRequest,
20 SerializedCredentialRequestLen, SerializedCredentialResponse, SerializedCredentialResponseLen,
21 SerializedIdentifier, SerializedIdentifiers,
22};
23use crate::opaque::MaskedResponseLen;
24use crate::serialization::{ConcatExt, SliceExt, UpdateExt};
25
26#[cfg_attr(
43 feature = "serde",
44 derive(serde::Deserialize, serde::Serialize),
45 serde(bound(deserialize = "'de: 'a", serialize = ""))
46)]
47#[derive_where(Clone, Debug, Eq, Hash, PartialEq, ZeroizeOnDrop)]
48pub struct Message<'a, CS: CipherSuite, KE: Group> {
49 pub(super) role: Role,
50 pub(super) context: SerializedContext<'a>,
51 pub(super) identifiers: SerializedIdentifiers<'a, KeGroup<CS>>,
52 pub(super) cache: CachedMessage<CS, KE>,
53}
54
55#[cfg_attr(
59 feature = "serde",
60 derive(serde::Deserialize, serde::Serialize),
61 serde(bound(deserialize = "'de: 'a", serialize = ""))
62)]
63#[derive_where(Clone, Debug, Eq, Hash, PartialEq, ZeroizeOnDrop)]
64pub struct VerifyMessage<'a, CS: CipherSuite, KE: Group> {
65 role: Role,
66 context: SerializedContext<'a>,
67 identifier: SerializedIdentifier<'a, KeGroup<CS>>,
68 pub(super) cache: CachedMessage<CS, KE>,
69}
70
71#[derive_where(Debug, Eq, Hash, PartialEq, ZeroizeOnDrop)]
76pub struct MessageBuilder<'a, CS: CipherSuite> {
77 pub(super) role: Role,
78 pub(super) context: SerializedContext<'a>,
79 pub(super) identifier: SerializedIdentifier<'a, KeGroup<CS>>,
80}
81
82#[cfg_attr(
91 feature = "serde",
92 derive(serde::Deserialize, serde::Serialize),
93 serde(bound = "")
94)]
95#[derive_where(Clone, Debug, Eq, Hash, PartialEq, Zeroize, ZeroizeOnDrop)]
96pub struct CachedMessage<CS: CipherSuite, KE: Group> {
97 pub(super) credential_request: SerializedCredentialRequest<CS>,
98 pub(super) ke1_message: Ke1MessageIter<KE>,
99 pub(super) credential_response: SerializedCredentialResponse<CS>,
100 pub(super) server_nonce: GenericArray<u8, NonceLen>,
101 pub(super) server_e_pk: GenericArray<u8, KE::PkLen>,
102 pub(super) server_mac: Output<KeHash<CS>>,
103}
104
105impl<CS: CipherSuite, KE: Group> Message<'_, CS, KE> {
106 pub fn sign_message(&self) -> impl Clone + Iterator<Item = &[u8]> {
108 self.context.iter().chain(self.post_message(Stage::Sign))
109 }
110
111 pub fn hash<KEH: Default + Clone + FixedOutput + Update>(&self) -> HashOutput<KEH> {
113 let mut context = KEH::default();
114 context.update_iter(self.context.iter());
115
116 let sign = context.clone().chain_iter(self.post_message(Stage::Sign));
117 let verify = context.chain_iter(self.post_message(Stage::Verify));
118
119 HashOutput { sign, verify }
120 }
121
122 fn post_message(&self, stage: Stage) -> impl Clone + Iterator<Item = &[u8]> {
123 let transcript = match (self.role, stage) {
124 (Role::Server, Stage::Sign) => Role::Server,
125 (Role::Server, Stage::Verify) => Role::Client,
126 (Role::Client, Stage::Sign) => Role::Client,
127 (Role::Client, Stage::Verify) => Role::Server,
128 };
129 let identifier = match transcript {
130 Role::Server => &self.identifiers.server,
131 Role::Client => &self.identifiers.client,
132 };
133
134 self.cache.post_message(transcript, identifier)
135 }
136
137 pub fn to_cached(&self) -> CachedMessage<CS, KE> {
143 self.cache.clone()
144 }
145}
146
147impl<CS: CipherSuite, KE: Group> VerifyMessage<'_, CS, KE> {
148 pub fn verify_message(&self) -> impl Clone + Iterator<Item = &[u8]> {
150 let transcript = match self.role {
151 Role::Server => Role::Client,
152 Role::Client => Role::Server,
153 };
154
155 self.context
156 .iter()
157 .chain(self.cache.post_message(transcript, &self.identifier))
158 }
159}
160
161impl<CS: CipherSuite, KE: Group> CachedMessage<CS, KE> {
162 fn post_message<'a>(
163 &'a self,
164 transcript: Role,
165 identifier: &'a SerializedIdentifier<'_, KeGroup<CS>>,
166 ) -> impl Clone + Iterator<Item = &'a [u8]> {
167 Some(identifier.iter())
168 .filter(|_| matches!(transcript, Role::Client))
169 .into_iter()
170 .flatten()
171 .chain(self.credential_request.iter())
172 .chain(self.ke1_message.iter())
173 .chain(
174 Some(identifier.iter())
175 .filter(|_| matches!(transcript, Role::Server))
176 .into_iter()
177 .flatten(),
178 )
179 .chain(self.credential_response.iter())
180 .chain([self.server_nonce.as_slice(), &self.server_e_pk])
181 .chain(Some(self.server_mac.as_slice()).filter(|_| matches!(transcript, Role::Client)))
182 }
183}
184
185#[cfg_attr(
186 feature = "serde",
187 derive(serde::Deserialize, serde::Serialize),
188 serde(bound = "")
189)]
190#[derive(Clone, Copy, Debug, Eq, Hash, PartialEq)]
191pub(super) enum Role {
192 Server,
193 Client,
194}
195
196impl Zeroize for Role {
197 fn zeroize(&mut self) {
198 *self = Self::Server;
199 }
200}
201
202enum Stage {
203 Sign,
204 Verify,
205}
206
207pub struct HashOutput<H> {
210 pub sign: H,
212 pub verify: H,
214}
215
216impl<'a, CS: CipherSuite> MessageBuilder<'a, CS> {
217 pub fn build<KE: Group>(self, cache: CachedMessage<CS, KE>) -> VerifyMessage<'a, CS, KE> {
223 VerifyMessage {
224 role: self.role,
225 context: self.context.clone(),
226 identifier: self.identifier.clone(),
227 cache,
228 }
229 }
230}
231
232impl<CS: CipherSuite, KE: Group> Deserialize for CachedMessage<CS, KE> {
233 fn deserialize_take(input: &mut &[u8]) -> Result<Self, ProtocolError> {
234 Ok(Self {
235 credential_request: SerializedCredentialRequest::deserialize_take(input)?,
236 ke1_message: Ke1MessageIter::deserialize_take(input)?,
237 credential_response: SerializedCredentialResponse::deserialize_take(input)?,
238 server_nonce: input.take_array("server nonce")?,
239 server_e_pk: input.take_array("serialized server ephemeral key")?,
240 server_mac: input.take_array("server mac")?.into_ha0_4(),
241 })
242 }
243}
244
245type CachedMessageLen<CS: CipherSuite, KE: Group> = Sum<
247 Sum<
248 Sum<
249 Sum<
250 Sum<SerializedCredentialRequestLen<CS>, Ke1MessageIterLen<KE>>,
251 SerializedCredentialResponseLen<CS>,
252 >,
253 NonceLen,
254 >,
255 KE::PkLen,
256 >,
257 OutputSize<KeHash<CS>>,
258>;
259
260impl<CS: CipherSuite, KE: Group> Serialize for CachedMessage<CS, KE>
261where
262 SerializedCredentialRequestLen<CS>: ArrayLength + Add<Ke1MessageIterLen<KE>>,
263 Sum<SerializedCredentialRequestLen<CS>, Ke1MessageIterLen<KE>>:
264 ArrayLength + Add<SerializedCredentialResponseLen<CS>>,
265 Sum<
266 Sum<SerializedCredentialRequestLen<CS>, Ke1MessageIterLen<KE>>,
267 SerializedCredentialResponseLen<CS>,
268 >: ArrayLength + Add<NonceLen>,
269 Sum<
270 Sum<
271 Sum<SerializedCredentialRequestLen<CS>, Ke1MessageIterLen<KE>>,
272 SerializedCredentialResponseLen<CS>,
273 >,
274 NonceLen,
275 >: ArrayLength + Add<KE::PkLen>,
276 Sum<
277 Sum<
278 Sum<
279 Sum<SerializedCredentialRequestLen<CS>, Ke1MessageIterLen<KE>>,
280 SerializedCredentialResponseLen<CS>,
281 >,
282 NonceLen,
283 >,
284 KE::PkLen,
285 >: ArrayLength + Add<OutputSize<KeHash<CS>>>,
286 CachedMessageLen<CS, KE>: ArrayLength,
287 NonceLen: Add<KE::PkLen>,
289 Ke1MessageIterLen<KE>: ArrayLength,
290 <OprfGroup<CS> as voprf::Group>::ElemLen: Add<NonceLen>,
292 Sum<<OprfGroup<CS> as voprf::Group>::ElemLen, NonceLen>:
293 ArrayLength + Add<MaskedResponseLen<CS>>,
294 SerializedCredentialResponseLen<CS>: ArrayLength,
295{
296 type Len = CachedMessageLen<CS, KE>;
297
298 fn serialize(&self) -> GenericArray<u8, Self::Len> {
299 self.credential_request
300 .serialize()
301 .cat(self.ke1_message.serialize())
302 .cat(self.credential_response.serialize())
303 .cat(self.server_nonce)
304 .cat(self.server_e_pk.clone())
305 .cat(GenericArray::from_slice(self.server_mac.as_slice()).clone())
306 }
307}