1use crate::coding::*;
5
6use super::{Message, Parameters, Version};
7
8const PARAM_PROBE: u64 = 0x1;
10const PARAM_PATH: u64 = 0x2;
12const PARAM_ROLE: u64 = 0x3;
14const PARAM_COST: u64 = 0x4;
16const PARAM_ORIGIN: u64 = 0x5;
18
19pub const DEFAULT_COST: u64 = 1;
26
27#[derive(Debug, Clone, Copy, PartialEq, Eq, PartialOrd, Ord, Default)]
33pub enum ProbeLevel {
34 #[default]
36 None,
37 Report,
39 Increase,
41}
42
43impl ProbeLevel {
44 fn from_code(code: u64) -> Self {
46 match code {
47 0 => Self::None,
48 1 => Self::Report,
49 _ => Self::Increase,
50 }
51 }
52
53 fn to_code(self) -> u64 {
55 match self {
56 Self::None => 0,
57 Self::Report => 1,
58 Self::Increase => 2,
59 }
60 }
61}
62
63#[derive(Debug, Clone, Copy, PartialEq, Eq)]
75#[non_exhaustive]
76pub enum Role {
77 Publisher,
79 Subscriber,
81}
82
83impl Role {
84 fn from_code(code: u64) -> Option<Self> {
89 match code {
90 1 => Some(Role::Publisher),
91 2 => Some(Role::Subscriber),
92 _ => None,
93 }
94 }
95
96 fn to_code(self) -> u64 {
98 match self {
99 Role::Publisher => 1,
100 Role::Subscriber => 2,
101 }
102 }
103
104 pub(crate) fn from_origins(publishes: bool, consumes: bool) -> Option<Self> {
109 match (publishes, consumes) {
110 (true, false) => Some(Role::Publisher),
111 (false, true) => Some(Role::Subscriber),
112 _ => None,
113 }
114 }
115
116 pub fn as_str(self) -> &'static str {
118 match self {
119 Role::Publisher => "publisher",
120 Role::Subscriber => "subscriber",
121 }
122 }
123}
124
125impl std::fmt::Display for Role {
126 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
127 f.write_str(self.as_str())
128 }
129}
130
131#[derive(Debug, Clone, Default, PartialEq, Eq)]
137pub struct Setup {
138 pub probe: ProbeLevel,
140 pub path: Option<String>,
147 pub role: Option<Role>,
151 pub cost: Option<u64>,
156 pub origin: Option<crate::Origin>,
162}
163
164impl Message for Setup {
165 fn decode_msg<R: bytes::Buf>(r: &mut R, version: Version) -> Result<Self, DecodeError> {
166 if !version.has_setup_stream() {
167 return Err(DecodeError::Version);
168 }
169
170 let params = Parameters::decode(r, version)?;
171 let probe = params
172 .get_varint(PARAM_PROBE)?
173 .map(ProbeLevel::from_code)
174 .unwrap_or_default();
175 let path = match params.get_bytes(PARAM_PATH) {
176 Some(bytes) => Some(
177 std::str::from_utf8(bytes)
178 .map_err(|_| DecodeError::InvalidValue)?
179 .to_string(),
180 ),
181 None => None,
182 };
183 let role = params.get_varint(PARAM_ROLE)?.and_then(Role::from_code);
184 let cost = params.get_varint(PARAM_COST)?;
185 let origin = params
188 .get_varint(PARAM_ORIGIN)?
189 .and_then(|id| crate::Origin::new(id).ok());
190
191 Ok(Self {
192 probe,
193 path,
194 role,
195 cost,
196 origin,
197 })
198 }
199
200 fn encode_msg<W: bytes::BufMut>(&self, w: &mut W, version: Version) -> Result<(), EncodeError> {
201 if !version.has_setup_stream() {
202 return Err(EncodeError::Version);
203 }
204
205 let mut params = Parameters::default();
206 if self.probe != ProbeLevel::None {
208 params.set_varint(PARAM_PROBE, self.probe.to_code());
209 }
210 if let Some(path) = &self.path {
211 params.set_bytes(PARAM_PATH, path.as_bytes().to_vec());
212 }
213 if let Some(role) = self.role {
216 params.set_varint(PARAM_ROLE, role.to_code());
217 }
218 if let Some(cost) = self.cost {
219 params.set_varint(PARAM_COST, cost);
220 }
221 if let Some(origin) = self.origin {
222 params.set_varint(PARAM_ORIGIN, origin.id());
223 }
224
225 params.encode(w, version)
226 }
227}
228
229#[derive(Clone, Default)]
235pub(crate) struct PeerSetup(kio::Shared<Option<Setup>>);
236
237impl PeerSetup {
238 pub fn set(&self, setup: Setup) {
240 *self.0.lock() = Some(setup);
241 }
242
243 pub async fn probe_level(&self) -> ProbeLevel {
245 self.wait(|setup| setup.probe).await
246 }
247
248 pub async fn cost(&self) -> Option<u64> {
251 self.wait(|setup| setup.cost).await
252 }
253
254 pub async fn origin(&self) -> Option<crate::Origin> {
257 self.wait(|setup| setup.origin).await
258 }
259
260 async fn wait<T>(&self, f: impl FnOnce(&Setup) -> T) -> T {
266 let slot = self
267 .0
268 .wait(|setup| {
269 if setup.is_some() {
270 std::task::Poll::Ready(())
271 } else {
272 std::task::Poll::Pending
273 }
274 })
275 .await;
276 f(slot.as_ref().expect("waited for Some"))
277 }
278}
279
280#[cfg(test)]
281mod tests {
282 use super::*;
283
284 fn round_trip(msg: &Setup) -> Setup {
285 let mut buf = bytes::BytesMut::new();
286 msg.encode(&mut buf, Version::Lite05).unwrap();
287 let mut slice = &buf[..];
288 let got = Setup::decode(&mut slice, Version::Lite05).unwrap();
289 assert!(bytes::Buf::remaining(&slice) == 0, "trailing bytes after decode");
290 got
291 }
292
293 #[test]
294 fn empty_round_trip() {
295 let msg = Setup::default();
296 assert_eq!(round_trip(&msg), msg);
297 }
298
299 #[test]
300 fn probe_levels_round_trip() {
301 for probe in [ProbeLevel::None, ProbeLevel::Report, ProbeLevel::Increase] {
302 let msg = Setup {
303 probe,
304 ..Default::default()
305 };
306 assert_eq!(round_trip(&msg), msg);
307 }
308 }
309
310 #[test]
311 fn cost_round_trip() {
312 for cost in [None, Some(0), Some(1), Some(7)] {
315 let msg = Setup {
316 cost,
317 ..Default::default()
318 };
319 assert_eq!(round_trip(&msg), msg);
320 }
321 }
322
323 #[test]
324 fn path_round_trip() {
325 let msg = Setup {
326 probe: ProbeLevel::Report,
327 path: Some("/room/123".to_string()),
328 ..Default::default()
329 };
330 assert_eq!(round_trip(&msg), msg);
331 }
332
333 #[test]
334 fn origin_round_trip() {
335 let msg = Setup {
336 origin: Some(crate::Origin::new(42).unwrap()),
337 ..Default::default()
338 };
339 assert_eq!(round_trip(&msg), msg);
340 }
341
342 #[test]
345 fn origin_zero_decodes_as_none() {
346 use crate::coding::Encode;
347
348 let version = Version::Lite05;
349 let mut params = Parameters::default();
350 params.set_varint(super::PARAM_ORIGIN, 0);
351 let mut body = bytes::BytesMut::new();
352 params.encode(&mut body, version).unwrap();
353 let mut buf = bytes::BytesMut::new();
355 (body.len() as u64).encode(&mut buf, version).unwrap();
356 buf.extend_from_slice(&body);
357 let mut slice = &buf[..];
358 let got = Setup::decode(&mut slice, version).unwrap();
359 assert_eq!(got.origin, None);
360 }
361
362 #[test]
363 fn empty_path_round_trips() {
364 let msg = Setup {
367 path: Some(String::new()),
368 ..Default::default()
369 };
370 assert_eq!(round_trip(&msg), msg);
371 }
372
373 #[test]
374 fn roles_round_trip() {
375 for role in [Some(Role::Publisher), Some(Role::Subscriber), None] {
376 let msg = Setup {
377 path: Some("/room/123".to_string()),
378 role,
379 ..Default::default()
380 };
381 assert_eq!(round_trip(&msg), msg);
382 }
383 }
384
385 #[test]
386 fn unknown_probe_level_saturates_to_increase() {
387 let mut params = Parameters::default();
390 params.set_varint(PARAM_PROBE, 99);
391 let mut body = Vec::new();
392 params.encode(&mut body, Version::Lite05).unwrap();
393
394 let mut buf = bytes::BytesMut::new();
395 body.len().encode(&mut buf, Version::Lite05).unwrap();
396 buf.extend_from_slice(&body);
397
398 let mut slice = &buf[..];
399 let got = Setup::decode(&mut slice, Version::Lite05).unwrap();
400 assert_eq!(got.probe, ProbeLevel::Increase);
401 }
402
403 #[test]
404 fn role_wire_codes() {
405 for (role, code) in [(Role::Publisher, 1u64), (Role::Subscriber, 2)] {
408 assert_eq!(role.to_code(), code);
409 assert_eq!(Role::from_code(code), Some(role));
410 }
411 }
412
413 #[test]
414 fn unknown_role_decodes_as_bidirectional() {
415 for code in [0u64, 9, 250] {
419 let mut params = Parameters::default();
420 params.set_varint(PARAM_ROLE, code);
421 let mut body = Vec::new();
422 params.encode(&mut body, Version::Lite05).unwrap();
423
424 let mut buf = bytes::BytesMut::new();
425 body.len().encode(&mut buf, Version::Lite05).unwrap();
426 buf.extend_from_slice(&body);
427
428 let mut slice = &buf[..];
429 let got = Setup::decode(&mut slice, Version::Lite05).unwrap();
430 assert_eq!(got.role, None, "role code {code} should decode as bidirectional");
431 }
432 }
433
434 #[test]
435 fn rejects_before_lite05() {
436 let msg = Setup::default();
437 let mut buf = bytes::BytesMut::new();
438 assert!(matches!(
439 msg.encode(&mut buf, Version::Lite04),
440 Err(EncodeError::Version)
441 ));
442 }
443
444 #[test]
445 fn ignores_unknown_parameters() {
446 let mut params = Parameters::default();
448 params.set_bytes(PARAM_PATH, b"/foo".to_vec());
449 params.set_bytes(0xbeef, b"whatever".to_vec());
450
451 let mut body = Vec::new();
452 params.encode(&mut body, Version::Lite05).unwrap();
453
454 let mut buf = bytes::BytesMut::new();
456 body.len().encode(&mut buf, Version::Lite05).unwrap();
457 buf.extend_from_slice(&body);
458
459 let mut slice = &buf[..];
460 let got = Setup::decode(&mut slice, Version::Lite05).unwrap();
461 assert_eq!(got.path.as_deref(), Some("/foo"));
462 }
463}