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>,
157 pub origin: Option<crate::Origin>,
163}
164
165impl Message for Setup {
166 fn decode_msg<R: bytes::Buf>(r: &mut R, version: Version) -> Result<Self, DecodeError> {
167 if !version.has_setup_stream() {
168 return Err(DecodeError::Version);
169 }
170
171 let params = Parameters::decode(r, version)?;
172 let probe = params
173 .get_varint(PARAM_PROBE)?
174 .map(ProbeLevel::from_code)
175 .unwrap_or_default();
176 let path = match params.get_bytes(PARAM_PATH) {
177 Some(bytes) => Some(
178 std::str::from_utf8(bytes)
179 .map_err(|_| DecodeError::InvalidValue)?
180 .to_string(),
181 ),
182 None => None,
183 };
184 let role = params.get_varint(PARAM_ROLE)?.and_then(Role::from_code);
185 let cost = params.get_varint(PARAM_COST)?;
186 let origin = params
189 .get_varint(PARAM_ORIGIN)?
190 .and_then(|id| crate::Origin::new(id).ok());
191
192 Ok(Self {
193 probe,
194 path,
195 role,
196 cost,
197 origin,
198 })
199 }
200
201 fn encode_msg<W: bytes::BufMut>(&self, w: &mut W, version: Version) -> Result<(), EncodeError> {
202 if !version.has_setup_stream() {
203 return Err(EncodeError::Version);
204 }
205
206 let mut params = Parameters::default();
207 if self.probe != ProbeLevel::None {
209 params.set_varint(PARAM_PROBE, self.probe.to_code());
210 }
211 if let Some(path) = &self.path {
212 params.set_bytes(PARAM_PATH, path.as_bytes().to_vec());
213 }
214 if let Some(role) = self.role {
217 params.set_varint(PARAM_ROLE, role.to_code());
218 }
219 if let Some(cost) = self.cost {
220 params.set_varint(PARAM_COST, cost);
221 }
222 if let Some(origin) = self.origin {
223 params.set_varint(PARAM_ORIGIN, origin.id());
224 }
225
226 params.encode(w, version)
227 }
228}
229
230#[derive(Clone, Default)]
236pub(crate) struct PeerSetup(kio::Shared<Option<Setup>>);
237
238impl PeerSetup {
239 pub fn set(&self, setup: Setup) {
241 *self.0.lock() = Some(setup);
242 }
243
244 pub async fn probe_level(&self) -> ProbeLevel {
246 self.wait(|setup| setup.probe).await
247 }
248
249 pub async fn cost(&self) -> Option<u64> {
252 self.wait(|setup| setup.cost).await
253 }
254
255 pub async fn origin(&self) -> Option<crate::Origin> {
258 self.wait(|setup| setup.origin).await
259 }
260
261 async fn wait<T>(&self, f: impl FnOnce(&Setup) -> T) -> T {
267 let slot = self
268 .0
269 .wait(|setup| {
270 if setup.is_some() {
271 std::task::Poll::Ready(())
272 } else {
273 std::task::Poll::Pending
274 }
275 })
276 .await;
277 f(slot.as_ref().expect("waited for Some"))
278 }
279}
280
281#[cfg(test)]
282mod tests {
283 use super::*;
284
285 fn round_trip(msg: &Setup) -> Setup {
286 let mut buf = bytes::BytesMut::new();
287 msg.encode(&mut buf, Version::Lite05).unwrap();
288 let mut slice = &buf[..];
289 let got = Setup::decode(&mut slice, Version::Lite05).unwrap();
290 assert!(bytes::Buf::remaining(&slice) == 0, "trailing bytes after decode");
291 got
292 }
293
294 #[test]
295 fn empty_round_trip() {
296 let msg = Setup::default();
297 assert_eq!(round_trip(&msg), msg);
298 }
299
300 #[test]
301 fn probe_levels_round_trip() {
302 for probe in [ProbeLevel::None, ProbeLevel::Report, ProbeLevel::Increase] {
303 let msg = Setup {
304 probe,
305 ..Default::default()
306 };
307 assert_eq!(round_trip(&msg), msg);
308 }
309 }
310
311 #[test]
312 fn cost_round_trip() {
313 for cost in [None, Some(0), Some(1), Some(7)] {
316 let msg = Setup {
317 cost,
318 ..Default::default()
319 };
320 assert_eq!(round_trip(&msg), msg);
321 }
322 }
323
324 #[test]
325 fn path_round_trip() {
326 let msg = Setup {
327 probe: ProbeLevel::Report,
328 path: Some("/room/123".to_string()),
329 ..Default::default()
330 };
331 assert_eq!(round_trip(&msg), msg);
332 }
333
334 #[test]
335 fn origin_round_trip() {
336 let msg = Setup {
337 origin: Some(crate::Origin::new(42).unwrap()),
338 ..Default::default()
339 };
340 assert_eq!(round_trip(&msg), msg);
341 }
342
343 #[test]
346 fn origin_zero_decodes_as_none() {
347 use crate::coding::Encode;
348
349 let version = Version::Lite05;
350 let mut params = Parameters::default();
351 params.set_varint(super::PARAM_ORIGIN, 0);
352 let mut body = bytes::BytesMut::new();
353 params.encode(&mut body, version).unwrap();
354 let mut buf = bytes::BytesMut::new();
356 (body.len() as u64).encode(&mut buf, version).unwrap();
357 buf.extend_from_slice(&body);
358 let mut slice = &buf[..];
359 let got = Setup::decode(&mut slice, version).unwrap();
360 assert_eq!(got.origin, None);
361 }
362
363 #[test]
364 fn empty_path_round_trips() {
365 let msg = Setup {
368 path: Some(String::new()),
369 ..Default::default()
370 };
371 assert_eq!(round_trip(&msg), msg);
372 }
373
374 #[test]
375 fn roles_round_trip() {
376 for role in [Some(Role::Publisher), Some(Role::Subscriber), None] {
377 let msg = Setup {
378 path: Some("/room/123".to_string()),
379 role,
380 ..Default::default()
381 };
382 assert_eq!(round_trip(&msg), msg);
383 }
384 }
385
386 #[test]
387 fn unknown_probe_level_saturates_to_increase() {
388 let mut params = Parameters::default();
391 params.set_varint(PARAM_PROBE, 99);
392 let mut body = Vec::new();
393 params.encode(&mut body, Version::Lite05).unwrap();
394
395 let mut buf = bytes::BytesMut::new();
396 body.len().encode(&mut buf, Version::Lite05).unwrap();
397 buf.extend_from_slice(&body);
398
399 let mut slice = &buf[..];
400 let got = Setup::decode(&mut slice, Version::Lite05).unwrap();
401 assert_eq!(got.probe, ProbeLevel::Increase);
402 }
403
404 #[test]
405 fn role_wire_codes() {
406 for (role, code) in [(Role::Publisher, 1u64), (Role::Subscriber, 2)] {
409 assert_eq!(role.to_code(), code);
410 assert_eq!(Role::from_code(code), Some(role));
411 }
412 }
413
414 #[test]
415 fn unknown_role_decodes_as_bidirectional() {
416 for code in [0u64, 9, 250] {
420 let mut params = Parameters::default();
421 params.set_varint(PARAM_ROLE, code);
422 let mut body = Vec::new();
423 params.encode(&mut body, Version::Lite05).unwrap();
424
425 let mut buf = bytes::BytesMut::new();
426 body.len().encode(&mut buf, Version::Lite05).unwrap();
427 buf.extend_from_slice(&body);
428
429 let mut slice = &buf[..];
430 let got = Setup::decode(&mut slice, Version::Lite05).unwrap();
431 assert_eq!(got.role, None, "role code {code} should decode as bidirectional");
432 }
433 }
434
435 #[test]
436 fn rejects_before_lite05() {
437 let msg = Setup::default();
438 let mut buf = bytes::BytesMut::new();
439 assert!(matches!(
440 msg.encode(&mut buf, Version::Lite04),
441 Err(EncodeError::Version)
442 ));
443 }
444
445 #[test]
446 fn ignores_unknown_parameters() {
447 let mut params = Parameters::default();
449 params.set_bytes(PARAM_PATH, b"/foo".to_vec());
450 params.set_bytes(0xbeef, b"whatever".to_vec());
451
452 let mut body = Vec::new();
453 params.encode(&mut body, Version::Lite05).unwrap();
454
455 let mut buf = bytes::BytesMut::new();
457 body.len().encode(&mut buf, Version::Lite05).unwrap();
458 buf.extend_from_slice(&body);
459
460 let mut slice = &buf[..];
461 let got = Setup::decode(&mut slice, Version::Lite05).unwrap();
462 assert_eq!(got.path.as_deref(), Some("/foo"));
463 }
464}