1use std::io::{Read, Write};
2
3use tls_codec::{Deserialize, DeserializeBytes, Serialize, Size, VLBytes};
4
5use crate::extensions::{Extension, ExtensionType, UnknownExtension};
6
7fn 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 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 let written = self.extension_type().tls_serialize(writer)?;
63
64 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 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 let extension_type = ExtensionType::tls_deserialize(bytes)?;
103 let extension_data = VLBytes::tls_deserialize(bytes)?;
104
105 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}