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                FrankenProposal::Custom(p) => p.tls_serialized_len(),
84            }
85    }
86}
87
88impl Serialize for FrankenProposal {
89    fn tls_serialize<W: std::io::Write>(&self, writer: &mut W) -> Result<usize, tls_codec::Error> {
90        let written = self.proposal_type().tls_serialize(writer)?;
91        match self {
92            FrankenProposal::Add(p) => p.tls_serialize(writer),
93            FrankenProposal::Update(p) => p.tls_serialize(writer),
94            FrankenProposal::Remove(p) => p.tls_serialize(writer),
95            FrankenProposal::PreSharedKey(p) => p.tls_serialize(writer),
96            FrankenProposal::ReInit(p) => p.tls_serialize(writer),
97            FrankenProposal::ExternalInit(p) => p.tls_serialize(writer),
98            FrankenProposal::GroupContextExtensions(p) => p.tls_serialize(writer),
99            #[cfg(feature = "extensions-draft")]
100            FrankenProposal::AppEphemeral(p) => p.tls_serialize(writer),
101            #[cfg(feature = "extensions-draft")]
102            FrankenProposal::AppDataUpdate(p) => p.tls_serialize(writer),
103            FrankenProposal::Custom(p) => p.payload.tls_serialize(writer),
104        }
105        .map(|l| written + l)
106    }
107}
108
109impl Deserialize for FrankenProposal {
110    fn tls_deserialize<R: std::io::Read>(bytes: &mut R) -> Result<Self, tls_codec::Error>
111    where
112        Self: Sized,
113    {
114        let proposal_type = FrankenProposalType::tls_deserialize(bytes)?;
115        let proposal = match proposal_type {
116            FrankenProposalType::Add => {
117                FrankenProposal::Add(FrankenAddProposal::tls_deserialize(bytes)?)
118            }
119            FrankenProposalType::Update => {
120                FrankenProposal::Update(FrankenUpdateProposal::tls_deserialize(bytes)?)
121            }
122            FrankenProposalType::Remove => {
123                FrankenProposal::Remove(FrankenRemoveProposal::tls_deserialize(bytes)?)
124            }
125            FrankenProposalType::PreSharedKey => {
126                FrankenProposal::PreSharedKey(FrankenPreSharedKeyProposal::tls_deserialize(bytes)?)
127            }
128            FrankenProposalType::Reinit => {
129                FrankenProposal::ReInit(FrankenReInitProposal::tls_deserialize(bytes)?)
130            }
131            FrankenProposalType::ExternalInit => {
132                FrankenProposal::ExternalInit(FrankenExternalInitProposal::tls_deserialize(bytes)?)
133            }
134            FrankenProposalType::GroupContextExtensions => FrankenProposal::GroupContextExtensions(
135                Vec::<FrankenExtension>::tls_deserialize(bytes)?,
136            ),
137            #[cfg(feature = "extensions-draft")]
138            FrankenProposalType::AppEphemeral => {
139                FrankenProposal::AppEphemeral(FrankenAppEphemeralProposal::tls_deserialize(bytes)?)
140            }
141            #[cfg(feature = "extensions-draft")]
142            FrankenProposalType::AppDataUpdate => FrankenProposal::AppDataUpdate(
143                FrankenAppDataUpdateProposal::tls_deserialize(bytes)?,
144            ),
145            FrankenProposalType::Custom(_) => {
146                let payload = VLBytes::tls_deserialize(bytes)?;
147                let custom_proposal = FrankenCustomProposal {
148                    proposal_type: proposal_type.into(),
149                    payload,
150                };
151                FrankenProposal::Custom(custom_proposal)
152            }
153        };
154        Ok(proposal)
155    }
156}
157
158impl DeserializeBytes for FrankenProposal {
159    fn tls_deserialize_bytes(bytes: &[u8]) -> Result<(Self, &[u8]), tls_codec::Error>
160    where
161        Self: Sized,
162    {
163        let mut bytes_ref = bytes;
164        let proposal = FrankenProposal::tls_deserialize(&mut bytes_ref)?;
165        Ok((proposal, bytes_ref))
166    }
167}
168
169impl Size for FrankenExtensionType {
170    fn tls_serialized_len(&self) -> usize {
171        2
172    }
173}
174
175impl Deserialize for FrankenExtensionType {
176    fn tls_deserialize<R: Read>(bytes: &mut R) -> Result<Self, Error>
177    where
178        Self: Sized,
179    {
180        let mut extension_type = [0u8; 2];
181        bytes.read_exact(&mut extension_type)?;
182
183        Ok(FrankenExtensionType::from(u16::from_be_bytes(
184            extension_type,
185        )))
186    }
187}
188
189impl DeserializeBytes for FrankenExtensionType {
190    fn tls_deserialize_bytes(bytes: &[u8]) -> Result<(Self, &[u8]), Error>
191    where
192        Self: Sized,
193    {
194        let mut bytes_ref = bytes;
195        let extension_type = FrankenExtensionType::tls_deserialize(&mut bytes_ref)?;
196        Ok((extension_type, bytes_ref))
197    }
198}
199
200impl Serialize for FrankenExtensionType {
201    fn tls_serialize<W: Write>(&self, writer: &mut W) -> Result<usize, Error> {
202        writer.write_all(&u16::from(*self).to_be_bytes())?;
203
204        Ok(2)
205    }
206}
207
208impl Size for FrankenExtension {
209    fn tls_serialized_len(&self) -> usize {
210        let extension_type_length = 2;
211        let extension_data_len = match self {
212            FrankenExtension::ApplicationId(e) => e.tls_serialized_len(),
213            FrankenExtension::RatchetTree(e) => e.tls_serialized_len(),
214            FrankenExtension::RequiredCapabilities(e) => e.tls_serialized_len(),
215            FrankenExtension::ExternalPub(e) => e.tls_serialized_len(),
216            FrankenExtension::ExternalSenders(e) => e.tls_serialized_len(),
217            FrankenExtension::LastResort => 0,
218            FrankenExtension::Unknown(_, e) => e.as_slice().len(),
219        };
220        let vlbytes_len_len = vlbytes_len_len(extension_data_len);
221        extension_type_length + vlbytes_len_len + extension_data_len
222    }
223}
224
225impl Serialize for FrankenExtension {
226    fn tls_serialize<W: Write>(&self, writer: &mut W) -> Result<usize, tls_codec::Error> {
227        let written = self.extension_type().tls_serialize(writer)?;
228
229        // subtract the two bytes for the type header
230        let extension_data_len = self.tls_serialized_len() - 2;
231        let mut extension_data = Vec::with_capacity(extension_data_len);
232
233        let _ = match self {
234            FrankenExtension::ApplicationId(e) => e.tls_serialize(&mut extension_data),
235            FrankenExtension::RatchetTree(e) => e.tls_serialize(&mut extension_data),
236            FrankenExtension::RequiredCapabilities(e) => e.tls_serialize(&mut extension_data),
237            FrankenExtension::ExternalPub(e) => e.tls_serialize(&mut extension_data),
238            FrankenExtension::ExternalSenders(e) => e.tls_serialize(&mut extension_data),
239            FrankenExtension::LastResort => Ok(0),
240            FrankenExtension::Unknown(_, e) => extension_data
241                .write_all(e.as_slice())
242                .map(|_| e.as_slice().len())
243                .map_err(|_| tls_codec::Error::EndOfStream),
244        }?;
245
246        Serialize::tls_serialize(&extension_data, writer).map(|l| l + written)
247    }
248}
249
250impl Deserialize for FrankenExtension {
251    fn tls_deserialize<R: Read>(bytes: &mut R) -> Result<Self, tls_codec::Error> {
252        // Read the extension type and extension data.
253        let extension_type = FrankenExtensionType::tls_deserialize(bytes)?;
254        let extension_data = VLBytes::tls_deserialize(bytes)?;
255
256        // Now deserialize the extension itself from the extension data.
257        let mut extension_data = extension_data.as_slice();
258        Ok(match extension_type {
259            FrankenExtensionType::ApplicationId => FrankenExtension::ApplicationId(
260                FrankenApplicationIdExtension::tls_deserialize(&mut extension_data)?,
261            ),
262            FrankenExtensionType::RatchetTree => FrankenExtension::RatchetTree(
263                FrankenRatchetTreeExtension::tls_deserialize(&mut extension_data)?,
264            ),
265            FrankenExtensionType::RequiredCapabilities => FrankenExtension::RequiredCapabilities(
266                FrankenRequiredCapabilitiesExtension::tls_deserialize(&mut extension_data)?,
267            ),
268            FrankenExtensionType::ExternalPub => FrankenExtension::ExternalPub(
269                FrankenExternalPubExtension::tls_deserialize(&mut extension_data)?,
270            ),
271            FrankenExtensionType::ExternalSenders => FrankenExtension::ExternalSenders(
272                FrankenExternalSendersExtension::tls_deserialize(&mut extension_data)?,
273            ),
274            FrankenExtensionType::LastResort => FrankenExtension::LastResort,
275            FrankenExtensionType::Unknown(unknown) => {
276                FrankenExtension::Unknown(unknown, extension_data.to_vec().into())
277            }
278        })
279    }
280}
281
282impl DeserializeBytes for FrankenExtension {
283    fn tls_deserialize_bytes(bytes: &[u8]) -> Result<(Self, &[u8]), tls_codec::Error>
284    where
285        Self: Sized,
286    {
287        let mut bytes_ref = bytes;
288        let extension = FrankenExtension::tls_deserialize(&mut bytes_ref)?;
289        Ok((extension, bytes_ref))
290    }
291}