1use bytes::Bytes;
6
7use crate::{
8 Version,
9 coding::{self, Decode, DecodeError, Encode, EncodeError, Sizer},
10 ietf, lite,
11};
12
13pub(crate) const MAX_SETUP_SIZE: usize = 64 * 1024;
15
16const CLIENT_SETUP: u8 = 0x20;
17const SERVER_SETUP: u8 = 0x21;
18
19pub(crate) const SETUP_V17: u64 = 0x2F00;
21
22#[derive(Debug, Clone, PartialEq, Eq)]
26pub struct Token {
27 pub kind: u64,
29 pub value: Vec<u8>,
31}
32
33impl Token {
34 pub const OUT_OF_BAND: u64 = 0x0;
36 pub const CAT: u64 = 0x1;
38}
39
40#[derive(Debug, Clone)]
42pub(crate) struct Setup {
43 pub parameters: Bytes,
44}
45
46impl Setup {
47 fn check_version(v: Version) {
48 match v {
49 Version::Ietf(ietf::Version::Draft14 | ietf::Version::Draft15 | ietf::Version::Draft16)
50 | Version::Lite(_) => unreachable!("Setup is draft-17+ only"),
51 _ => {}
52 }
53 }
54}
55
56impl Encode<Version> for Setup {
57 fn encode<W: bytes::BufMut>(&self, w: &mut W, v: Version) -> Result<(), EncodeError> {
58 Self::check_version(v);
59 SETUP_V17.encode(w, v)?;
60 u16::try_from(self.parameters.len())
61 .map_err(|_| EncodeError::TooLarge)?
62 .encode(w, v)?;
63 if w.remaining_mut() < self.parameters.len() {
64 return Err(EncodeError::Short);
65 }
66 w.put_slice(&self.parameters);
67 Ok(())
68 }
69}
70
71impl Decode<Version> for Setup {
72 fn decode<R: bytes::Buf>(r: &mut R, v: Version) -> Result<Self, DecodeError> {
73 Self::check_version(v);
74 let kind = u64::decode(r, v)?;
75 if kind != SETUP_V17 {
76 return Err(DecodeError::InvalidValue);
77 }
78 let size = u16::decode(r, v)? as usize;
80 if r.remaining() < size {
81 return Err(DecodeError::Short);
82 }
83 let msg = r.copy_to_bytes(size);
84 Ok(Self { parameters: msg })
85 }
86}
87
88#[derive(Clone, Copy, Debug, PartialEq, Eq)]
89enum SetupVersion {
90 Draft14,
91 Draft15Plus,
92 Modern,
94 LiteLegacy,
95 Unsupported,
96}
97
98impl SetupVersion {
99 fn from_version(v: Version) -> Self {
100 match v {
101 Version::Ietf(ietf::Version::Draft14) => Self::Draft14,
102 Version::Ietf(ietf::Version::Draft15) | Version::Ietf(ietf::Version::Draft16) => Self::Draft15Plus,
103 Version::Ietf(ietf::Version::Draft17)
104 | Version::Ietf(ietf::Version::Draft18)
105 | Version::Ietf(ietf::Version::Draft19)
106 | Version::Ietf(ietf::Version::Draft20)
107 | Version::Ietf(ietf::Version::Draft21)
108 | Version::Ietf(ietf::Version::Draft22) => Self::Modern,
109 Version::Lite(lite::Version::Lite01) | Version::Lite(lite::Version::Lite02) => Self::LiteLegacy,
110 Version::Lite(_) => Self::Unsupported,
111 }
112 }
113}
114
115#[derive(Debug, Clone)]
117pub(crate) struct Client {
118 pub versions: coding::Versions,
120
121 pub parameters: Bytes,
123}
124
125impl Client {
126 fn encode_inner<W: bytes::BufMut>(&self, w: &mut W, v: Version) -> Result<(), EncodeError> {
127 match SetupVersion::from_version(v) {
128 SetupVersion::Draft15Plus => {
129 }
131 SetupVersion::Draft14 | SetupVersion::LiteLegacy => self.versions.encode(w, v)?,
132 SetupVersion::Modern | SetupVersion::Unsupported => return Err(EncodeError::Version),
133 };
134 if w.remaining_mut() < self.parameters.len() {
135 return Err(EncodeError::Short);
136 }
137 w.put_slice(&self.parameters);
138 Ok(())
139 }
140}
141
142impl Decode<Version> for Client {
143 fn decode<R: bytes::Buf>(r: &mut R, v: Version) -> Result<Self, DecodeError> {
145 let kind = u8::decode(r, v)?;
146 if kind != CLIENT_SETUP {
147 return Err(DecodeError::InvalidValue);
148 }
149
150 let size = match SetupVersion::from_version(v) {
151 SetupVersion::Draft14 | SetupVersion::Draft15Plus => u16::decode(r, v)? as usize,
152 SetupVersion::LiteLegacy => usize::decode(r, v)?,
153 SetupVersion::Modern | SetupVersion::Unsupported => return Err(DecodeError::Version),
154 };
155
156 if size > MAX_SETUP_SIZE {
157 return Err(DecodeError::MessageTooLarge {
158 size,
159 max: MAX_SETUP_SIZE,
160 });
161 }
162
163 if r.remaining() < size {
164 return Err(DecodeError::Short);
165 }
166
167 let mut msg = r.copy_to_bytes(size);
168
169 let versions = match SetupVersion::from_version(v) {
170 SetupVersion::Draft15Plus => {
171 coding::Versions::from([v.into()])
173 }
174 SetupVersion::Draft14 | SetupVersion::LiteLegacy => {
175 coding::Versions::decode(&mut msg, v).map_err(DecodeError::complete)?
176 }
177 SetupVersion::Modern | SetupVersion::Unsupported => return Err(DecodeError::Version),
178 };
179
180 Ok(Self {
181 versions,
182 parameters: msg,
183 })
184 }
185}
186
187impl Encode<Version> for Client {
188 fn encode<W: bytes::BufMut>(&self, w: &mut W, v: Version) -> Result<(), EncodeError> {
190 CLIENT_SETUP.encode(w, v)?;
191
192 let mut sizer = Sizer::default();
193 self.encode_inner(&mut sizer, v)?;
194 let size = sizer.size;
195 if size > MAX_SETUP_SIZE {
196 return Err(EncodeError::TooLarge);
197 }
198
199 match SetupVersion::from_version(v) {
200 SetupVersion::Draft14 | SetupVersion::Draft15Plus => {
201 u16::try_from(size).map_err(|_| EncodeError::TooLarge)?.encode(w, v)?;
202 }
203 SetupVersion::LiteLegacy => (size as u64).encode(w, v)?,
204 SetupVersion::Modern | SetupVersion::Unsupported => return Err(EncodeError::Version),
205 }
206 self.encode_inner(w, v)
207 }
208}
209
210#[derive(Debug, Clone)]
212pub(crate) struct Server {
213 pub version: coding::Version,
215
216 pub parameters: Bytes,
218}
219
220impl Server {
221 fn encode_inner<W: bytes::BufMut>(&self, w: &mut W, v: Version) -> Result<(), EncodeError> {
222 match SetupVersion::from_version(v) {
223 SetupVersion::Draft15Plus => {
224 }
226 SetupVersion::Draft14 | SetupVersion::LiteLegacy => self.version.encode(w, v)?,
227 SetupVersion::Modern | SetupVersion::Unsupported => return Err(EncodeError::Version),
228 };
229 if w.remaining_mut() < self.parameters.len() {
230 return Err(EncodeError::Short);
231 }
232 w.put_slice(&self.parameters);
233 Ok(())
234 }
235}
236
237impl Encode<Version> for Server {
238 fn encode<W: bytes::BufMut>(&self, w: &mut W, v: Version) -> Result<(), EncodeError> {
240 SERVER_SETUP.encode(w, v)?;
241
242 let mut sizer = Sizer::default();
243 self.encode_inner(&mut sizer, v)?;
244 let size = sizer.size;
245 if size > MAX_SETUP_SIZE {
246 return Err(EncodeError::TooLarge);
247 }
248
249 match SetupVersion::from_version(v) {
250 SetupVersion::Draft14 | SetupVersion::Draft15Plus => {
251 u16::try_from(size).map_err(|_| EncodeError::TooLarge)?.encode(w, v)?;
252 }
253 SetupVersion::LiteLegacy => (size as u64).encode(w, v)?,
254 SetupVersion::Modern | SetupVersion::Unsupported => return Err(EncodeError::Version),
255 }
256
257 self.encode_inner(w, v)
258 }
259}
260
261impl Decode<Version> for Server {
262 fn decode<R: bytes::Buf>(r: &mut R, v: Version) -> Result<Self, DecodeError> {
264 let kind = u8::decode(r, v)?;
265 if kind != SERVER_SETUP {
266 return Err(DecodeError::InvalidValue);
267 }
268
269 let size = match SetupVersion::from_version(v) {
270 SetupVersion::Draft14 | SetupVersion::Draft15Plus => u16::decode(r, v)? as usize,
271 SetupVersion::LiteLegacy => usize::decode(r, v)?,
272 SetupVersion::Modern | SetupVersion::Unsupported => return Err(DecodeError::Version),
273 };
274
275 if size > MAX_SETUP_SIZE {
276 return Err(DecodeError::MessageTooLarge {
277 size,
278 max: MAX_SETUP_SIZE,
279 });
280 }
281
282 if r.remaining() < size {
283 return Err(DecodeError::Short);
284 }
285
286 let mut msg = r.copy_to_bytes(size);
287 let version = match SetupVersion::from_version(v) {
288 SetupVersion::Draft15Plus => v.into(),
289 SetupVersion::Draft14 | SetupVersion::LiteLegacy => {
290 coding::Version::decode(&mut msg, v).map_err(DecodeError::complete)?
291 }
292 SetupVersion::Modern | SetupVersion::Unsupported => return Err(DecodeError::Version),
293 };
294
295 Ok(Self {
296 version,
297 parameters: msg,
298 })
299 }
300}
301
302#[cfg(test)]
303mod tests {
304 use super::*;
305
306 #[test]
308 fn encode_enforces_the_setup_limit() {
309 let v = Version::Lite(lite::Version::Lite01);
310 let parameters = Bytes::from(vec![0; MAX_SETUP_SIZE]);
311
312 let client = Client {
313 versions: coding::Versions::from([v.into()]),
314 parameters: parameters.clone(),
315 };
316 let mut buf = Vec::new();
317 assert!(matches!(client.encode(&mut buf, v), Err(EncodeError::TooLarge)));
318
319 let server = Server {
320 version: v.into(),
321 parameters,
322 };
323 let mut buf = Vec::new();
324 assert!(matches!(server.encode(&mut buf, v), Err(EncodeError::TooLarge)));
325 }
326}