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