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.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 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 let extension_type = FrankenExtensionType::tls_deserialize(bytes)?;
256 let extension_data = VLBytes::tls_deserialize(bytes)?;
257
258 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 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 #[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}