Skip to main content

opaque_vx/key_exchange/sigma_i/
message.rs

1// SPDX-License-Identifier: MIT OR Apache-2.0
2// Copyright (c) VexaHub and contributors.
3// Copyright (c) Meta Platforms, Inc. and affiliates.
4
5use 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/// This holds the message to be signed and the message to be verified.
27///
28/// If your signature protocol requires pre-hashes, you can call [`hash()`].
29///
30/// If you require the actual message, call [`sign_message()`]. To get the
31/// message to verify, call [`to_cached()`] to create a [`CachedMessage`] and
32/// save it in [`SignatureProtocol::VerifyState`], which you can then use in
33/// [`SignatureProtocol::verify()`] with [`MessageBuilder`] to create
34/// [`VerifyMessage`].
35///
36/// [`hash()`]: super::Message::hash
37/// [`sign_message()`]: super::Message::sign_message
38/// [`to_cached()`]: super::Message::to_cached
39/// [`SignatureProtocol::sign()`]: super::SignatureProtocol::sign
40/// [`SignatureProtocol::verify()`]: super::SignatureProtocol::verify
41/// [`SignatureProtocol::VerifyState`]: super::SignatureProtocol::VerifyState
42#[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/// This holds the message to be verified.
56///
57/// Create it by using [`MessageBuilder::build()`] with [`CachedMessage`].
58#[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/// Used to build [`VerifyMessage`]. It is only available in
72/// [`SignatureProtocol::verify()`].
73///
74/// [`SignatureProtocol::verify()`]: super::SignatureProtocol::verify
75#[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/// Created by [`Message::to_cached()`]. This is used to save the message to be
83/// verified in [`SignatureProtocol::VerifyState`].
84///
85/// Use [`MessageBuilder::build()`] to create [`VerifyMessage`] in
86/// [`SignatureProtocol::verify()`].
87///
88/// [`SignatureProtocol::verify()`]: super::SignatureProtocol::verify
89/// [`SignatureProtocol::VerifyState`]: super::SignatureProtocol::VerifyState
90#[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    /// Returns the message to be signed.
107    pub fn sign_message(&self) -> impl Clone + Iterator<Item = &[u8]> {
108        self.context.iter().chain(self.post_message(Stage::Sign))
109    }
110
111    /// Returns the hash of both messages.
112    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    /// Create a [`CachedMessage`], which can be saved in
138    /// [`SignatureProtocol::VerifyState`] and create a [`VerifyMessage`] with
139    /// [`MessageBuilder::build()`].
140    ///
141    /// [`SignatureProtocol::VerifyState`]: super::SignatureProtocol::VerifyState
142    pub fn to_cached(&self) -> CachedMessage<CS, KE> {
143        self.cache.clone()
144    }
145}
146
147impl<CS: CipherSuite, KE: Group> VerifyMessage<'_, CS, KE> {
148    /// Returns the message to be verified.
149    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
207/// Returned by [`Message::hash()`] containing the hash of the message to be
208/// signed and the message to be verified.
209pub struct HashOutput<H> {
210    /// The hash of the message to be signed.
211    pub sign: H,
212    /// The hash of the message to be verified.
213    pub verify: H,
214}
215
216impl<'a, CS: CipherSuite> MessageBuilder<'a, CS> {
217    /// Creates a [`VerifyMessage`]. [`CachedMessage`] can be created by
218    /// [`Message::to_cached()`] and stored in
219    /// [`SignatureProtocol::VerifyState`].
220    ///
221    /// [`SignatureProtocol::VerifyState`]: super::SignatureProtocol::VerifyState
222    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
245/// Length of [`CachedMessage`].
246type 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    // Ke1MessageIter
288    NonceLen: Add<KE::PkLen>,
289    Ke1MessageIterLen<KE>: ArrayLength,
290    // CredentialResponseParts
291    <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}