Skip to main content

openmls/extensions/
codec.rs

1use std::io::{Read, Write};
2
3use tls_codec::{Deserialize, DeserializeBytes, Serialize, Size, VLBytes};
4
5use crate::extensions::{Extension, ExtensionType, UnknownExtension};
6
7/// Known extension types must consume the bounded payload without silently dropping trailing bytes.
8fn deserialize_extension_exact<T: Deserialize>(
9    extension_data: &[u8],
10) -> Result<T, tls_codec::Error> {
11    T::tls_deserialize_exact(extension_data)
12}
13
14fn vlbytes_len_len(length: usize) -> usize {
15    if length <= 0x3f {
16        1
17    } else if length <= 0x3fff {
18        2
19    } else if length <= 0x3fff_ffff {
20        4
21    } else {
22        8
23    }
24}
25
26impl Size for Extension {
27    #[inline]
28    fn tls_serialized_len(&self) -> usize {
29        let extension_type_length = 2;
30
31        // We truncate here and don't catch errors for anything that's
32        // too long.
33        // This will be caught when (de)serializing.
34        let extension_data_len = match self {
35            Extension::ApplicationId(e) => e.tls_serialized_len(),
36            Extension::RatchetTree(e) => e.tls_serialized_len(),
37            Extension::RequiredCapabilities(e) => e.tls_serialized_len(),
38            Extension::ExternalPub(e) => e.tls_serialized_len(),
39            Extension::ExternalSenders(e) => e.tls_serialized_len(),
40            Extension::LastResort(e) => e.tls_serialized_len(),
41            #[cfg(feature = "extensions-draft")]
42            Extension::AppDataDictionary(e) => e.tls_serialized_len(),
43            Extension::Unknown(_, e) => e.0.len(),
44        };
45
46        let vlbytes_len_len = vlbytes_len_len(extension_data_len);
47
48        extension_type_length + vlbytes_len_len + extension_data_len
49    }
50}
51
52impl Size for &Extension {
53    #[inline]
54    fn tls_serialized_len(&self) -> usize {
55        Extension::tls_serialized_len(*self)
56    }
57}
58
59impl Serialize for Extension {
60    fn tls_serialize<W: Write>(&self, writer: &mut W) -> Result<usize, tls_codec::Error> {
61        // First write the extension type.
62        let written = self.extension_type().tls_serialize(writer)?;
63
64        // Now serialize the extension into a separate byte vector.
65        let extension_data_len = self.tls_serialized_len();
66        let mut extension_data = Vec::with_capacity(extension_data_len);
67
68        let extension_data_written = match self {
69            Extension::ApplicationId(e) => e.tls_serialize(&mut extension_data),
70            Extension::RatchetTree(e) => e.tls_serialize(&mut extension_data),
71            Extension::RequiredCapabilities(e) => e.tls_serialize(&mut extension_data),
72            Extension::ExternalPub(e) => e.tls_serialize(&mut extension_data),
73            Extension::ExternalSenders(e) => e.tls_serialize(&mut extension_data),
74            #[cfg(feature = "extensions-draft")]
75            Extension::AppDataDictionary(e) => e.tls_serialize(&mut extension_data),
76            Extension::LastResort(e) => e.tls_serialize(&mut extension_data),
77            Extension::Unknown(_, e) => extension_data
78                .write_all(e.0.as_slice())
79                .map(|_| e.0.len())
80                .map_err(|_| tls_codec::Error::EndOfStream),
81        }?;
82        debug_assert_eq!(
83            extension_data_written,
84            extension_data_len - 2 - vlbytes_len_len(extension_data_written)
85        );
86        debug_assert_eq!(extension_data_written, extension_data.len());
87
88        // Write the serialized extension out.
89        extension_data.tls_serialize(writer).map(|l| l + written)
90    }
91}
92
93impl Serialize for &Extension {
94    fn tls_serialize<W: Write>(&self, writer: &mut W) -> Result<usize, tls_codec::Error> {
95        Extension::tls_serialize(*self, writer)
96    }
97}
98
99impl Deserialize for Extension {
100    fn tls_deserialize<R: Read>(bytes: &mut R) -> Result<Self, tls_codec::Error> {
101        // Read the extension type and extension data.
102        let extension_type = ExtensionType::tls_deserialize(bytes)?;
103        let extension_data = VLBytes::tls_deserialize(bytes)?;
104
105        // Now deserialize the extension itself from the extension data.
106        let extension_data = extension_data.as_slice();
107        Ok(match extension_type {
108            ExtensionType::ApplicationId => {
109                Extension::ApplicationId(deserialize_extension_exact(extension_data)?)
110            }
111            ExtensionType::RatchetTree => {
112                Extension::RatchetTree(deserialize_extension_exact(extension_data)?)
113            }
114            ExtensionType::RequiredCapabilities => {
115                Extension::RequiredCapabilities(deserialize_extension_exact(extension_data)?)
116            }
117            ExtensionType::ExternalPub => {
118                Extension::ExternalPub(deserialize_extension_exact(extension_data)?)
119            }
120            ExtensionType::ExternalSenders => {
121                Extension::ExternalSenders(deserialize_extension_exact(extension_data)?)
122            }
123            #[cfg(feature = "extensions-draft")]
124            ExtensionType::AppDataDictionary => {
125                Extension::AppDataDictionary(deserialize_extension_exact(extension_data)?)
126            }
127            ExtensionType::LastResort => {
128                Extension::LastResort(deserialize_extension_exact(extension_data)?)
129            }
130            ExtensionType::Grease(grease) | ExtensionType::Unknown(grease) => {
131                Extension::Unknown(grease, UnknownExtension(extension_data.to_vec()))
132            }
133        })
134    }
135}
136
137impl DeserializeBytes for Extension {
138    fn tls_deserialize_bytes(bytes: &[u8]) -> Result<(Self, &[u8]), tls_codec::Error>
139    where
140        Self: Sized,
141    {
142        let mut bytes_ref = bytes;
143        let extension = Extension::tls_deserialize(&mut bytes_ref)?;
144        Ok((extension, bytes_ref))
145    }
146}
147
148#[cfg(test)]
149mod tests {
150    use super::*;
151    #[cfg(feature = "extensions-draft")]
152    use crate::extensions::AppDataDictionaryExtension;
153    use crate::{
154        credentials::CredentialType,
155        extensions::{
156            ApplicationIdExtension, ExternalPubExtension, ExternalSendersExtension,
157            LastResortExtension, RequiredCapabilitiesExtension,
158        },
159        messages::proposals::ProposalType,
160        treesync::RatchetTreeIn,
161    };
162
163    fn serialize_extension(extension_type: ExtensionType, payload: Vec<u8>) -> Vec<u8> {
164        let mut serialized = extension_type.tls_serialize_detached().unwrap();
165        serialized.extend(VLBytes::from(payload).tls_serialize_detached().unwrap());
166        serialized
167    }
168
169    fn known_extensions() -> Vec<(ExtensionType, Vec<u8>)> {
170        let extensions = vec![
171            (
172                ExtensionType::ApplicationId,
173                ApplicationIdExtension::new(&[1, 2, 3])
174                    .tls_serialize_detached()
175                    .unwrap(),
176            ),
177            (
178                ExtensionType::RatchetTree,
179                RatchetTreeIn::from_nodes(vec![])
180                    .tls_serialize_detached()
181                    .unwrap(),
182            ),
183            (
184                ExtensionType::RequiredCapabilities,
185                RequiredCapabilitiesExtension::new(
186                    &[ExtensionType::ApplicationId],
187                    &[ProposalType::Add],
188                    &[CredentialType::Basic],
189                )
190                .tls_serialize_detached()
191                .unwrap(),
192            ),
193            (
194                ExtensionType::ExternalPub,
195                ExternalPubExtension::new(vec![4, 5, 6].into())
196                    .tls_serialize_detached()
197                    .unwrap(),
198            ),
199            (
200                ExtensionType::ExternalSenders,
201                ExternalSendersExtension::new()
202                    .tls_serialize_detached()
203                    .unwrap(),
204            ),
205            (
206                ExtensionType::LastResort,
207                LastResortExtension::new().tls_serialize_detached().unwrap(),
208            ),
209        ];
210
211        #[cfg(feature = "extensions-draft")]
212        let extensions = {
213            let mut extensions = extensions;
214            extensions.push((
215                ExtensionType::AppDataDictionary,
216                AppDataDictionaryExtension::default()
217                    .tls_serialize_detached()
218                    .unwrap(),
219            ));
220            extensions
221        };
222
223        extensions
224    }
225
226    #[test]
227    fn known_extensions_round_trip() {
228        for (extension_type, payload) in known_extensions() {
229            let serialized = serialize_extension(extension_type, payload);
230            let extension = Extension::tls_deserialize_exact(&serialized).unwrap();
231            assert_eq!(extension.tls_serialize_detached().unwrap(), serialized);
232        }
233    }
234
235    #[test]
236    fn known_extensions_reject_trailing_payload_bytes() {
237        for (extension_type, mut payload) in known_extensions() {
238            payload.extend([0xa5, 0x5a]);
239            let serialized = serialize_extension(extension_type, payload);
240            assert_eq!(
241                Extension::tls_deserialize_exact(&serialized).unwrap_err(),
242                tls_codec::Error::TrailingData
243            );
244        }
245    }
246
247    #[cfg(feature = "extensions-draft")]
248    #[test]
249    fn app_data_dictionary_uses_exact_payload_decoding() {
250        use crate::extensions::AppDataDictionary;
251
252        let mut dictionary = AppDataDictionary::new();
253        dictionary.insert(0x8001, vec![1, 2, 3]);
254        let mut payload = AppDataDictionaryExtension::new(dictionary)
255            .tls_serialize_detached()
256            .unwrap();
257        payload.push(0xff);
258
259        let serialized = serialize_extension(ExtensionType::AppDataDictionary, payload);
260        assert_eq!(
261            Extension::tls_deserialize_exact(serialized).unwrap_err(),
262            tls_codec::Error::TrailingData
263        );
264    }
265
266    #[test]
267    fn opaque_extensions_round_trip_arbitrary_payload() {
268        let payload = vec![0x00, 0xff, 0x01, 0xfe, 0x80];
269        for extension_type in [
270            ExtensionType::Unknown(0xf042),
271            ExtensionType::Grease(0x0a0a),
272        ] {
273            let serialized = serialize_extension(extension_type, payload.clone());
274            let extension = Extension::tls_deserialize_exact(&serialized).unwrap();
275            assert_eq!(
276                extension,
277                Extension::Unknown(u16::from(extension_type), UnknownExtension(payload.clone()))
278            );
279            assert_eq!(extension.tls_serialize_detached().unwrap(), serialized);
280        }
281    }
282}