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_HOP: 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,
41 Increase,
43}
44
45impl ProbeLevel {
46 pub fn detect<S: crate::transport::poll::Session>(session: &S) -> Self {
58 use web_transport_trait::Stats as _;
59 let stats = session.stats();
60 match stats.estimated_send_rate().is_some() || stats.rtt().is_some() {
61 true => Self::Report,
62 false => Self::None,
63 }
64 }
65
66 fn from_code(code: u64) -> Self {
68 match code {
69 0 => Self::None,
70 1 => Self::Report,
71 _ => Self::Increase,
72 }
73 }
74
75 fn to_code(self) -> u64 {
77 match self {
78 Self::None => 0,
79 Self::Report => 1,
80 Self::Increase => 2,
81 }
82 }
83}
84
85#[derive(Debug, Clone, Copy, PartialEq, Eq)]
97#[non_exhaustive]
98pub enum Role {
99 Publisher,
101 Subscriber,
103}
104
105impl Role {
106 fn from_code(code: u64) -> Option<Self> {
111 match code {
112 1 => Some(Role::Publisher),
113 2 => Some(Role::Subscriber),
114 _ => None,
115 }
116 }
117
118 fn to_code(self) -> u64 {
120 match self {
121 Role::Publisher => 1,
122 Role::Subscriber => 2,
123 }
124 }
125
126 pub(crate) fn from_origins(publishes: bool, consumes: bool) -> Option<Self> {
131 match (publishes, consumes) {
132 (true, false) => Some(Role::Publisher),
133 (false, true) => Some(Role::Subscriber),
134 _ => None,
135 }
136 }
137
138 pub fn as_str(self) -> &'static str {
140 match self {
141 Role::Publisher => "publisher",
142 Role::Subscriber => "subscriber",
143 }
144 }
145}
146
147impl std::fmt::Display for Role {
148 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
149 f.write_str(self.as_str())
150 }
151}
152
153#[derive(Debug, Clone, Default, PartialEq, Eq)]
159pub struct Setup {
160 pub probe: ProbeLevel,
162 pub path: Option<String>,
169 pub role: Option<Role>,
173 pub cost: Option<u64>,
179 pub hop: Option<crate::Hop>,
185}
186
187impl Message for Setup {
188 const MAX_SIZE: usize = crate::setup::MAX_SETUP_SIZE;
189
190 fn decode_msg<R: bytes::Buf>(r: &mut R, version: Version) -> Result<Self, DecodeError> {
191 if !version.has_setup_stream() {
192 return Err(DecodeError::Version);
193 }
194
195 let params = Parameters::decode(r, version)?;
196 let probe = params
197 .get_varint(PARAM_PROBE, version)?
198 .map(ProbeLevel::from_code)
199 .unwrap_or_default();
200 let path = match params.get_bytes(PARAM_PATH) {
201 Some(bytes) => Some(
202 std::str::from_utf8(bytes)
203 .map_err(|_| DecodeError::InvalidValue)?
204 .to_string(),
205 ),
206 None => None,
207 };
208 let role = params.get_varint(PARAM_ROLE, version)?.and_then(Role::from_code);
209 let cost = params.get_varint(PARAM_COST, version)?;
210 let hop = params
213 .get_varint(PARAM_HOP, version)?
214 .and_then(|id| crate::Hop::new(id).ok());
215
216 Ok(Self {
217 probe,
218 path,
219 role,
220 cost,
221 hop,
222 })
223 }
224
225 fn encode_msg<W: bytes::BufMut>(&self, w: &mut W, version: Version) -> Result<(), EncodeError> {
226 if !version.has_setup_stream() {
227 return Err(EncodeError::Version);
228 }
229
230 let mut params = Parameters::default();
231 if self.probe != ProbeLevel::None {
233 params.set_varint(PARAM_PROBE, self.probe.to_code(), version)?;
234 }
235 if let Some(path) = &self.path {
236 params.set_bytes(PARAM_PATH, path.as_bytes().to_vec());
237 }
238 if let Some(role) = self.role {
241 params.set_varint(PARAM_ROLE, role.to_code(), version)?;
242 }
243 if let Some(cost) = self.cost {
244 params.set_varint(PARAM_COST, cost, version)?;
245 }
246 if let Some(hop) = self.hop {
247 params.set_varint(PARAM_HOP, hop.id(), version)?;
248 }
249
250 params.encode(w, version)
251 }
252}
253
254#[derive(Clone, Default)]
260pub(crate) struct PeerSetup(kio::Shared<Option<Setup>>);
261
262impl PeerSetup {
263 pub fn set(&self, setup: Setup) {
265 *self.0.lock() = Some(setup);
266 }
267
268 pub fn poll_probe_level(&self, waiter: &kio::Waiter) -> std::task::Poll<ProbeLevel> {
270 self.poll_get(waiter, |setup| setup.probe)
271 }
272
273 pub fn poll_cost(&self, waiter: &kio::Waiter) -> std::task::Poll<Option<u64>> {
276 self.poll_get(waiter, |setup| setup.cost)
277 }
278
279 pub fn poll_hop(&self, waiter: &kio::Waiter) -> std::task::Poll<Option<crate::Hop>> {
282 self.poll_get(waiter, |setup| setup.hop)
283 }
284
285 fn poll_get<T>(&self, waiter: &kio::Waiter, f: impl FnOnce(&Setup) -> T) -> std::task::Poll<T> {
291 let slot = std::task::ready!(self.0.poll(waiter, |setup| {
292 if setup.is_some() {
293 std::task::Poll::Ready(())
294 } else {
295 std::task::Poll::Pending
296 }
297 }));
298 std::task::Poll::Ready(f(slot.as_ref().expect("waited for Some")))
299 }
300}
301
302#[cfg(test)]
303mod tests {
304 use super::*;
305
306 fn round_trip(msg: &Setup) -> Setup {
307 let mut buf = bytes::BytesMut::new();
308 msg.encode(&mut buf, Version::Lite05).unwrap();
309 let mut slice = &buf[..];
310 let got = Setup::decode(&mut slice, Version::Lite05).unwrap();
311 assert!(bytes::Buf::remaining(&slice) == 0, "trailing bytes after decode");
312 got
313 }
314
315 #[test]
316 fn empty_round_trip() {
317 let msg = Setup::default();
318 assert_eq!(round_trip(&msg), msg);
319 }
320
321 #[test]
325 fn detect_reports_nothing_without_stats() {
326 use crate::lite::test_transport::{SinkSession, SinkStats};
327 let session = SinkSession::new(Default::default()).with_stats(SinkStats::default());
328 assert_eq!(ProbeLevel::detect(&session), ProbeLevel::None);
329 }
330
331 #[test]
334 fn detect_reports_with_either_metric() {
335 use crate::lite::test_transport::{SinkSession, SinkStats};
336
337 let rtt_only = SinkStats::default().with_rtt(std::time::Duration::from_millis(40));
338 let session = SinkSession::new(Default::default()).with_stats(rtt_only);
339 assert_eq!(ProbeLevel::detect(&session), ProbeLevel::Report);
340
341 let rate_only = SinkStats::default().with_send_rate(1_000_000);
342 let session = SinkSession::new(Default::default()).with_stats(rate_only);
343 assert_eq!(ProbeLevel::detect(&session), ProbeLevel::Report);
344 }
345
346 #[test]
347 fn probe_levels_round_trip() {
348 for probe in [ProbeLevel::None, ProbeLevel::Report, ProbeLevel::Increase] {
349 let msg = Setup {
350 probe,
351 ..Default::default()
352 };
353 assert_eq!(round_trip(&msg), msg);
354 }
355 }
356
357 #[test]
358 fn cost_round_trip() {
359 for cost in [None, Some(0), Some(1), Some(7)] {
362 let msg = Setup {
363 cost,
364 ..Default::default()
365 };
366 assert_eq!(round_trip(&msg), msg);
367 }
368 }
369
370 #[test]
373 fn parameter_values_use_the_version_codec() {
374 let msg = Setup {
375 cost: Some(100),
376 ..Default::default()
377 };
378 for (version, wire) in [
379 (Version::Lite06, &[0x05, 0x01, 0x04, 0x02, 0x40, 0x64][..]),
380 (Version::Lite07, &[0x04, 0x01, 0x04, 0x01, 0x64][..]),
381 ] {
382 let mut buf = Vec::new();
383 msg.encode(&mut buf, version).unwrap();
384 assert_eq!(buf, wire, "{version}");
385 assert_eq!(Setup::decode(&mut &buf[..], version).unwrap(), msg, "{version}");
386 }
387 }
388
389 #[test]
391 fn encode_enforces_the_setup_limit() {
392 let at_limit = crate::setup::MAX_SETUP_SIZE - 6;
394 let msg = Setup {
395 path: Some("a".repeat(at_limit)),
396 ..Default::default()
397 };
398 assert_eq!(round_trip(&msg), msg);
399
400 let msg = Setup {
401 path: Some("a".repeat(at_limit + 1)),
402 ..Default::default()
403 };
404 let mut buf = bytes::BytesMut::new();
405 assert!(matches!(
406 msg.encode(&mut buf, Version::Lite05),
407 Err(EncodeError::TooLarge)
408 ));
409 assert!(buf.is_empty());
410 }
411
412 #[test]
413 fn path_round_trip() {
414 let msg = Setup {
415 probe: ProbeLevel::Report,
416 path: Some("/room/123".to_string()),
417 ..Default::default()
418 };
419 assert_eq!(round_trip(&msg), msg);
420 }
421
422 #[test]
423 fn hop_round_trip() {
424 let msg = Setup {
425 hop: Some(crate::Hop::new(42).unwrap()),
426 ..Default::default()
427 };
428 assert_eq!(round_trip(&msg), msg);
429 }
430
431 #[test]
434 fn hop_zero_decodes_as_none() {
435 use crate::coding::Encode;
436
437 let version = Version::Lite05;
438 let mut params = Parameters::default();
439 params.set_varint(super::PARAM_HOP, 0, version).unwrap();
440 let mut body = bytes::BytesMut::new();
441 params.encode(&mut body, version).unwrap();
442 let mut buf = bytes::BytesMut::new();
444 (body.len() as u64).encode(&mut buf, version).unwrap();
445 buf.extend_from_slice(&body);
446 let mut slice = &buf[..];
447 let got = Setup::decode(&mut slice, version).unwrap();
448 assert_eq!(got.hop, None);
449 }
450
451 #[test]
452 fn empty_path_round_trips() {
453 let msg = Setup {
456 path: Some(String::new()),
457 ..Default::default()
458 };
459 assert_eq!(round_trip(&msg), msg);
460 }
461
462 #[test]
463 fn roles_round_trip() {
464 for role in [Some(Role::Publisher), Some(Role::Subscriber), None] {
465 let msg = Setup {
466 path: Some("/room/123".to_string()),
467 role,
468 ..Default::default()
469 };
470 assert_eq!(round_trip(&msg), msg);
471 }
472 }
473
474 #[test]
475 fn unknown_probe_level_saturates_to_increase() {
476 let mut params = Parameters::default();
479 params.set_varint(PARAM_PROBE, 99, Version::Lite05).unwrap();
480 let mut body = Vec::new();
481 params.encode(&mut body, Version::Lite05).unwrap();
482
483 let mut buf = bytes::BytesMut::new();
484 body.len().encode(&mut buf, Version::Lite05).unwrap();
485 buf.extend_from_slice(&body);
486
487 let mut slice = &buf[..];
488 let got = Setup::decode(&mut slice, Version::Lite05).unwrap();
489 assert_eq!(got.probe, ProbeLevel::Increase);
490 }
491
492 #[test]
493 fn role_wire_codes() {
494 for (role, code) in [(Role::Publisher, 1u64), (Role::Subscriber, 2)] {
497 assert_eq!(role.to_code(), code);
498 assert_eq!(Role::from_code(code), Some(role));
499 }
500 }
501
502 #[test]
503 fn unknown_role_decodes_as_bidirectional() {
504 for code in [0u64, 9, 250] {
508 let mut params = Parameters::default();
509 params.set_varint(PARAM_ROLE, code, Version::Lite05).unwrap();
510 let mut body = Vec::new();
511 params.encode(&mut body, Version::Lite05).unwrap();
512
513 let mut buf = bytes::BytesMut::new();
514 body.len().encode(&mut buf, Version::Lite05).unwrap();
515 buf.extend_from_slice(&body);
516
517 let mut slice = &buf[..];
518 let got = Setup::decode(&mut slice, Version::Lite05).unwrap();
519 assert_eq!(got.role, None, "role code {code} should decode as bidirectional");
520 }
521 }
522
523 #[test]
524 fn rejects_before_lite05() {
525 let msg = Setup::default();
526 let mut buf = bytes::BytesMut::new();
527 assert!(matches!(
528 msg.encode(&mut buf, Version::Lite04),
529 Err(EncodeError::Version)
530 ));
531 }
532
533 #[test]
534 fn ignores_unknown_parameters() {
535 let mut params = Parameters::default();
537 params.set_bytes(PARAM_PATH, b"/foo".to_vec());
538 params.set_bytes(0xbeef, b"whatever".to_vec());
539
540 let mut body = Vec::new();
541 params.encode(&mut body, Version::Lite05).unwrap();
542
543 let mut buf = bytes::BytesMut::new();
545 body.len().encode(&mut buf, Version::Lite05).unwrap();
546 buf.extend_from_slice(&body);
547
548 let mut slice = &buf[..];
549 let got = Setup::decode(&mut slice, Version::Lite05).unwrap();
550 assert_eq!(got.path.as_deref(), Some("/foo"));
551 }
552}