Skip to main content

openmls/test_utils/frankenstein/
codec.rs

1use std::io::{Read, Write};
2
3use tls_codec::*;
4
5use super::{
6    extensions::{
7        FrankenApplicationIdExtension, FrankenExtension, FrankenExtensionType,
8        FrankenExternalPubExtension, FrankenExternalSendersExtension, FrankenRatchetTreeExtension,
9        FrankenRequiredCapabilitiesExtension,
10    },
11    FrankenAddProposal, FrankenCustomProposal, FrankenExternalInitProposal,
12    FrankenPreSharedKeyProposal, FrankenProposal, FrankenProposalType, FrankenReInitProposal,
13    FrankenRemoveProposal, FrankenUpdateProposal,
14};
15
16#[cfg(feature = "extensions-draft")]
17use super::{FrankenAppDataUpdateProposal, FrankenAppEphemeralProposal};
18
19fn vlbytes_len_len(length: usize) -> usize {
20    if length <= 0x3f {
21        1
22    } else if length <= 0x3fff {
23        2
24    } else if length <= 0x3fff_ffff {
25        4
26    } else {
27        8
28    }
29}
30
31impl Size for FrankenProposalType {
32    fn tls_serialized_len(&self) -> usize {
33        2
34    }
35}
36
37impl Deserialize for FrankenProposalType {
38    fn tls_deserialize<R: Read>(bytes: &mut R) -> Result<Self, Error>
39    where
40        Self: Sized,
41    {
42        let mut proposal_type = [0u8; 2];
43        bytes.read_exact(&mut proposal_type)?;
44
45        Ok(FrankenProposalType::from(u16::from_be_bytes(proposal_type)))
46    }
47}
48
49impl Serialize for FrankenProposalType {
50    fn tls_serialize<W: Write>(&self, writer: &mut W) -> Result<usize, Error> {
51        writer.write_all(&u16::from(*self).to_be_bytes())?;
52
53        Ok(2)
54    }
55}
56
57impl DeserializeBytes for FrankenProposalType {
58    fn tls_deserialize_bytes(bytes: &[u8]) -> Result<(Self, &[u8]), Error>
59    where
60        Self: Sized,
61    {
62        let mut bytes_ref = bytes;
63        let proposal_type = FrankenProposalType::tls_deserialize(&mut bytes_ref)?;
64        Ok((proposal_type, bytes_ref))
65    }
66}
67
68impl Size for FrankenProposal {
69    fn tls_serialized_len(&self) -> usize {
70        self.proposal_type().tls_serialized_len()
71            + match self {
72                FrankenProposal::Add(p) => p.tls_serialized_len(),
73                FrankenProposal::Update(p) => p.tls_serialized_len(),
74                FrankenProposal::Remove(p) => p.tls_serialized_len(),
75                FrankenProposal::PreSharedKey(p) => p.tls_serialized_len(),
76                FrankenProposal::ReInit(p) => p.tls_serialized_len(),
77                FrankenProposal::ExternalInit(p) => p.tls_serialized_len(),
78                FrankenProposal::GroupContextExtensions(p) => p.tls_serialized_len(),
79                #[cfg(feature = "extensions-draft")]
80                FrankenProposal::AppEphemeral(p) => p.tls_serialized_len(),
81                #[cfg(feature = "extensions-draft")]
82                FrankenProposal::AppDataUpdate(p) => p.tls_serialized_len(),
83                // Only the payload is written; the proposal type is already
84                // accounted for above.
85                FrankenProposal::Custom(p) => p.payload.tls_serialized_len(),
86            }
87    }
88}
89
90impl Serialize for FrankenProposal {
91    fn tls_serialize<W: std::io::Write>(&self, writer: &mut W) -> Result<usize, tls_codec::Error> {
92        let written = self.proposal_type().tls_serialize(writer)?;
93        match self {
94            FrankenProposal::Add(p) => p.tls_serialize(writer),
95            FrankenProposal::Update(p) => p.tls_serialize(writer),
96            FrankenProposal::Remove(p) => p.tls_serialize(writer),
97            FrankenProposal::PreSharedKey(p) => p.tls_serialize(writer),
98            FrankenProposal::ReInit(p) => p.tls_serialize(writer),
99            FrankenProposal::ExternalInit(p) => p.tls_serialize(writer),
100            FrankenProposal::GroupContextExtensions(p) => p.tls_serialize(writer),
101            #[cfg(feature = "extensions-draft")]
102            FrankenProposal::AppEphemeral(p) => p.tls_serialize(writer),
103            #[cfg(feature = "extensions-draft")]
104            FrankenProposal::AppDataUpdate(p) => p.tls_serialize(writer),
105            FrankenProposal::Custom(p) => p.payload.tls_serialize(writer),
106        }
107        .map(|l| written + l)
108    }
109}
110
111impl Deserialize for FrankenProposal {
112    fn tls_deserialize<R: std::io::Read>(bytes: &mut R) -> Result<Self, tls_codec::Error>
113    where
114        Self: Sized,
115    {
116        let proposal_type = FrankenProposalType::tls_deserialize(bytes)?;
117        let proposal = match proposal_type {
118            FrankenProposalType::Add => {
119                FrankenProposal::Add(FrankenAddProposal::tls_deserialize(bytes)?)
120            }
121            FrankenProposalType::Update => {
122                FrankenProposal::Update(FrankenUpdateProposal::tls_deserialize(bytes)?)
123            }
124            FrankenProposalType::Remove => {
125                FrankenProposal::Remove(FrankenRemoveProposal::tls_deserialize(bytes)?)
126            }
127            FrankenProposalType::PreSharedKey => {
128                FrankenProposal::PreSharedKey(FrankenPreSharedKeyProposal::tls_deserialize(bytes)?)
129            }
130            FrankenProposalType::Reinit => {
131                FrankenProposal::ReInit(FrankenReInitProposal::tls_deserialize(bytes)?)
132            }
133            FrankenProposalType::ExternalInit => {
134                FrankenProposal::ExternalInit(FrankenExternalInitProposal::tls_deserialize(bytes)?)
135            }
136            FrankenProposalType::GroupContextExtensions => FrankenProposal::GroupContextExtensions(
137                Vec::<FrankenExtension>::tls_deserialize(bytes)?,
138            ),
139            #[cfg(feature = "extensions-draft")]
140            FrankenProposalType::AppEphemeral => {
141                FrankenProposal::AppEphemeral(FrankenAppEphemeralProposal::tls_deserialize(bytes)?)
142            }
143            #[cfg(feature = "extensions-draft")]
144            FrankenProposalType::AppDataUpdate => FrankenProposal::AppDataUpdate(
145                FrankenAppDataUpdateProposal::tls_deserialize(bytes)?,
146            ),
147            FrankenProposalType::Custom(_) => {
148                let payload = VLBytes::tls_deserialize(bytes)?;
149                let custom_proposal = FrankenCustomProposal {
150                    proposal_type: proposal_type.into(),
151                    payload,
152                };
153                FrankenProposal::Custom(custom_proposal)
154            }
155        };
156        Ok(proposal)
157    }
158}
159
160impl DeserializeBytes for FrankenProposal {
161    fn tls_deserialize_bytes(bytes: &[u8]) -> Result<(Self, &[u8]), tls_codec::Error>
162    where
163        Self: Sized,
164    {
165        let mut bytes_ref = bytes;
166        let proposal = FrankenProposal::tls_deserialize(&mut bytes_ref)?;
167        Ok((proposal, bytes_ref))
168    }
169}
170
171impl Size for FrankenExtensionType {
172    fn tls_serialized_len(&self) -> usize {
173        2
174    }
175}
176
177impl Deserialize for FrankenExtensionType {
178    fn tls_deserialize<R: Read>(bytes: &mut R) -> Result<Self, Error>
179    where
180        Self: Sized,
181    {
182        let mut extension_type = [0u8; 2];
183        bytes.read_exact(&mut extension_type)?;
184
185        Ok(FrankenExtensionType::from(u16::from_be_bytes(
186            extension_type,
187        )))
188    }
189}
190
191impl DeserializeBytes for FrankenExtensionType {
192    fn tls_deserialize_bytes(bytes: &[u8]) -> Result<(Self, &[u8]), Error>
193    where
194        Self: Sized,
195    {
196        let mut bytes_ref = bytes;
197        let extension_type = FrankenExtensionType::tls_deserialize(&mut bytes_ref)?;
198        Ok((extension_type, bytes_ref))
199    }
200}
201
202impl Serialize for FrankenExtensionType {
203    fn tls_serialize<W: Write>(&self, writer: &mut W) -> Result<usize, Error> {
204        writer.write_all(&u16::from(*self).to_be_bytes())?;
205
206        Ok(2)
207    }
208}
209
210impl Size for FrankenExtension {
211    fn tls_serialized_len(&self) -> usize {
212        let extension_type_length = 2;
213        let extension_data_len = match self {
214            FrankenExtension::ApplicationId(e) => e.tls_serialized_len(),
215            FrankenExtension::RatchetTree(e) => e.tls_serialized_len(),
216            FrankenExtension::RequiredCapabilities(e) => e.tls_serialized_len(),
217            FrankenExtension::ExternalPub(e) => e.tls_serialized_len(),
218            FrankenExtension::ExternalSenders(e) => e.tls_serialized_len(),
219            FrankenExtension::LastResort => 0,
220            FrankenExtension::Unknown(_, e) => e.as_slice().len(),
221        };
222        let vlbytes_len_len = vlbytes_len_len(extension_data_len);
223        extension_type_length + vlbytes_len_len + extension_data_len
224    }
225}
226
227impl Serialize for FrankenExtension {
228    fn tls_serialize<W: Write>(&self, writer: &mut W) -> Result<usize, tls_codec::Error> {
229        let written = self.extension_type().tls_serialize(writer)?;
230
231        // subtract the two bytes for the type header
232        let extension_data_len = self.tls_serialized_len() - 2;
233        let mut extension_data = Vec::with_capacity(extension_data_len);
234
235        let _ = match self {
236            FrankenExtension::ApplicationId(e) => e.tls_serialize(&mut extension_data),
237            FrankenExtension::RatchetTree(e) => e.tls_serialize(&mut extension_data),
238            FrankenExtension::RequiredCapabilities(e) => e.tls_serialize(&mut extension_data),
239            FrankenExtension::ExternalPub(e) => e.tls_serialize(&mut extension_data),
240            FrankenExtension::ExternalSenders(e) => e.tls_serialize(&mut extension_data),
241            FrankenExtension::LastResort => Ok(0),
242            FrankenExtension::Unknown(_, e) => extension_data
243                .write_all(e.as_slice())
244                .map(|_| e.as_slice().len())
245                .map_err(|_| tls_codec::Error::EndOfStream),
246        }?;
247
248        Serialize::tls_serialize(&extension_data, writer).map(|l| l + written)
249    }
250}
251
252impl Deserialize for FrankenExtension {
253    fn tls_deserialize<R: Read>(bytes: &mut R) -> Result<Self, tls_codec::Error> {
254        // Read the extension type and extension data.
255        let extension_type = FrankenExtensionType::tls_deserialize(bytes)?;
256        let extension_data = VLBytes::tls_deserialize(bytes)?;
257
258        // Now deserialize the extension itself from the extension data.
259        let mut extension_data = extension_data.as_slice();
260        Ok(match extension_type {
261            FrankenExtensionType::ApplicationId => FrankenExtension::ApplicationId(
262                FrankenApplicationIdExtension::tls_deserialize(&mut extension_data)?,
263            ),
264            FrankenExtensionType::RatchetTree => FrankenExtension::RatchetTree(
265                FrankenRatchetTreeExtension::tls_deserialize(&mut extension_data)?,
266            ),
267            FrankenExtensionType::RequiredCapabilities => FrankenExtension::RequiredCapabilities(
268                FrankenRequiredCapabilitiesExtension::tls_deserialize(&mut extension_data)?,
269            ),
270            FrankenExtensionType::ExternalPub => FrankenExtension::ExternalPub(
271                FrankenExternalPubExtension::tls_deserialize(&mut extension_data)?,
272            ),
273            FrankenExtensionType::ExternalSenders => FrankenExtension::ExternalSenders(
274                FrankenExternalSendersExtension::tls_deserialize(&mut extension_data)?,
275            ),
276            FrankenExtensionType::LastResort => FrankenExtension::LastResort,
277            FrankenExtensionType::Unknown(unknown) => {
278                FrankenExtension::Unknown(unknown, extension_data.to_vec().into())
279            }
280        })
281    }
282}
283
284impl DeserializeBytes for FrankenExtension {
285    fn tls_deserialize_bytes(bytes: &[u8]) -> Result<(Self, &[u8]), tls_codec::Error>
286    where
287        Self: Sized,
288    {
289        let mut bytes_ref = bytes;
290        let extension = FrankenExtension::tls_deserialize(&mut bytes_ref)?;
291        Ok((extension, bytes_ref))
292    }
293}
294
295#[cfg(test)]
296mod tests {
297    use super::*;
298    use crate::test_utils::frankenstein::{
299        FrankenCommit, FrankenCustomProposal, FrankenExternalPsk, FrankenPreSharedKeyId,
300        FrankenPreSharedKeyProposal, FrankenProposalOrRef, FrankenPsk, FrankenReInitProposal,
301        FrankenRemoveProposal,
302    };
303
304    #[cfg(feature = "extensions-draft")]
305    use crate::messages::proposals::AppDataUpdateOperation;
306
307    /// Proposals that can be built without a key package or a leaf node.
308    fn proposals() -> Vec<FrankenProposal> {
309        let proposals = vec![
310            FrankenProposal::Remove(FrankenRemoveProposal { removed: 3 }),
311            FrankenProposal::PreSharedKey(FrankenPreSharedKeyProposal {
312                psk: FrankenPreSharedKeyId {
313                    psk: FrankenPsk::External(FrankenExternalPsk {
314                        psk_id: vec![7, 8].into(),
315                    }),
316                    psk_nonce: vec![9; 32].into(),
317                },
318            }),
319            FrankenProposal::ReInit(FrankenReInitProposal {
320                group_id: vec![1, 2, 3].into(),
321                version: 1,
322                ciphersuite: 1,
323                extensions: vec![FrankenExtension::LastResort],
324            }),
325            FrankenProposal::ExternalInit(FrankenExternalInitProposal {
326                kem_output: vec![4; 32].into(),
327            }),
328            FrankenProposal::GroupContextExtensions(vec![FrankenExtension::Unknown(
329                0xf042,
330                vec![5, 6].into(),
331            )]),
332            FrankenProposal::Custom(FrankenCustomProposal {
333                proposal_type: 0xf001,
334                payload: vec![1, 2, 3, 4].into(),
335            }),
336        ];
337
338        #[cfg(feature = "extensions-draft")]
339        let proposals = {
340            let mut proposals = proposals;
341            proposals.push(FrankenProposal::AppEphemeral(FrankenAppEphemeralProposal {
342                component_id: 0x8001,
343                data: vec![1, 2, 3].into(),
344            }));
345            proposals.push(FrankenProposal::AppDataUpdate(
346                FrankenAppDataUpdateProposal {
347                    component_id: 0x8001,
348                    operation: AppDataUpdateOperation::Update(vec![4, 5, 6].into()),
349                },
350            ));
351            proposals
352        };
353
354        proposals
355    }
356
357    #[test]
358    fn proposal_length_matches_serialization() {
359        for proposal in proposals() {
360            let serialized = proposal.tls_serialize_detached().unwrap();
361            assert_eq!(
362                proposal.tls_serialized_len(),
363                serialized.len(),
364                "wrong length for {proposal:?}"
365            );
366            assert_eq!(
367                FrankenProposal::tls_deserialize_exact(&serialized).unwrap(),
368                proposal
369            );
370        }
371    }
372
373    /// A wrong length breaks any enclosing struct that writes a length prefix,
374    /// so exercise a proposal nested in a commit as well.
375    #[test]
376    fn commit_with_a_custom_proposal_round_trips() {
377        let commit = FrankenCommit {
378            proposals: vec![FrankenProposalOrRef::Proposal(FrankenProposal::Custom(
379                FrankenCustomProposal {
380                    proposal_type: 0xf001,
381                    payload: vec![1, 2, 3, 4].into(),
382                },
383            ))],
384            path: None,
385        };
386
387        let serialized = commit.tls_serialize_detached().unwrap();
388        assert_eq!(
389            FrankenCommit::tls_deserialize_exact(&serialized).unwrap(),
390            commit
391        );
392    }
393}