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 fn decode_msg<R: bytes::Buf>(r: &mut R, version: Version) -> Result<Self, DecodeError> {
189 if !version.has_setup_stream() {
190 return Err(DecodeError::Version);
191 }
192
193 let params = Parameters::decode(r, version)?;
194 let probe = params
195 .get_varint(PARAM_PROBE)?
196 .map(ProbeLevel::from_code)
197 .unwrap_or_default();
198 let path = match params.get_bytes(PARAM_PATH) {
199 Some(bytes) => Some(
200 std::str::from_utf8(bytes)
201 .map_err(|_| DecodeError::InvalidValue)?
202 .to_string(),
203 ),
204 None => None,
205 };
206 let role = params.get_varint(PARAM_ROLE)?.and_then(Role::from_code);
207 let cost = params.get_varint(PARAM_COST)?;
208 let hop = params.get_varint(PARAM_HOP)?.and_then(|id| crate::Hop::new(id).ok());
211
212 Ok(Self {
213 probe,
214 path,
215 role,
216 cost,
217 hop,
218 })
219 }
220
221 fn encode_msg<W: bytes::BufMut>(&self, w: &mut W, version: Version) -> Result<(), EncodeError> {
222 if !version.has_setup_stream() {
223 return Err(EncodeError::Version);
224 }
225
226 let mut params = Parameters::default();
227 if self.probe != ProbeLevel::None {
229 params.set_varint(PARAM_PROBE, self.probe.to_code());
230 }
231 if let Some(path) = &self.path {
232 params.set_bytes(PARAM_PATH, path.as_bytes().to_vec());
233 }
234 if let Some(role) = self.role {
237 params.set_varint(PARAM_ROLE, role.to_code());
238 }
239 if let Some(cost) = self.cost {
240 params.set_varint(PARAM_COST, cost);
241 }
242 if let Some(hop) = self.hop {
243 params.set_varint(PARAM_HOP, hop.id());
244 }
245
246 params.encode(w, version)
247 }
248}
249
250#[derive(Clone, Default)]
256pub(crate) struct PeerSetup(kio::Shared<Option<Setup>>);
257
258impl PeerSetup {
259 pub fn set(&self, setup: Setup) {
261 *self.0.lock() = Some(setup);
262 }
263
264 pub fn poll_probe_level(&self, waiter: &kio::Waiter) -> std::task::Poll<ProbeLevel> {
266 self.poll_get(waiter, |setup| setup.probe)
267 }
268
269 pub fn poll_cost(&self, waiter: &kio::Waiter) -> std::task::Poll<Option<u64>> {
272 self.poll_get(waiter, |setup| setup.cost)
273 }
274
275 pub fn poll_hop(&self, waiter: &kio::Waiter) -> std::task::Poll<Option<crate::Hop>> {
278 self.poll_get(waiter, |setup| setup.hop)
279 }
280
281 fn poll_get<T>(&self, waiter: &kio::Waiter, f: impl FnOnce(&Setup) -> T) -> std::task::Poll<T> {
287 let slot = std::task::ready!(self.0.poll(waiter, |setup| {
288 if setup.is_some() {
289 std::task::Poll::Ready(())
290 } else {
291 std::task::Poll::Pending
292 }
293 }));
294 std::task::Poll::Ready(f(slot.as_ref().expect("waited for Some")))
295 }
296}
297
298#[cfg(test)]
299mod tests {
300 use super::*;
301
302 fn round_trip(msg: &Setup) -> Setup {
303 let mut buf = bytes::BytesMut::new();
304 msg.encode(&mut buf, Version::Lite05).unwrap();
305 let mut slice = &buf[..];
306 let got = Setup::decode(&mut slice, Version::Lite05).unwrap();
307 assert!(bytes::Buf::remaining(&slice) == 0, "trailing bytes after decode");
308 got
309 }
310
311 #[test]
312 fn empty_round_trip() {
313 let msg = Setup::default();
314 assert_eq!(round_trip(&msg), msg);
315 }
316
317 #[test]
321 fn detect_reports_nothing_without_stats() {
322 use crate::lite::test_transport::{SinkSession, SinkStats};
323 let session = SinkSession::new(Default::default()).with_stats(SinkStats::default());
324 assert_eq!(ProbeLevel::detect(&session), ProbeLevel::None);
325 }
326
327 #[test]
330 fn detect_reports_with_either_metric() {
331 use crate::lite::test_transport::{SinkSession, SinkStats};
332
333 let rtt_only = SinkStats::default().with_rtt(std::time::Duration::from_millis(40));
334 let session = SinkSession::new(Default::default()).with_stats(rtt_only);
335 assert_eq!(ProbeLevel::detect(&session), ProbeLevel::Report);
336
337 let rate_only = SinkStats::default().with_send_rate(1_000_000);
338 let session = SinkSession::new(Default::default()).with_stats(rate_only);
339 assert_eq!(ProbeLevel::detect(&session), ProbeLevel::Report);
340 }
341
342 #[test]
343 fn probe_levels_round_trip() {
344 for probe in [ProbeLevel::None, ProbeLevel::Report, ProbeLevel::Increase] {
345 let msg = Setup {
346 probe,
347 ..Default::default()
348 };
349 assert_eq!(round_trip(&msg), msg);
350 }
351 }
352
353 #[test]
354 fn cost_round_trip() {
355 for cost in [None, Some(0), Some(1), Some(7)] {
358 let msg = Setup {
359 cost,
360 ..Default::default()
361 };
362 assert_eq!(round_trip(&msg), msg);
363 }
364 }
365
366 #[test]
367 fn path_round_trip() {
368 let msg = Setup {
369 probe: ProbeLevel::Report,
370 path: Some("/room/123".to_string()),
371 ..Default::default()
372 };
373 assert_eq!(round_trip(&msg), msg);
374 }
375
376 #[test]
377 fn hop_round_trip() {
378 let msg = Setup {
379 hop: Some(crate::Hop::new(42).unwrap()),
380 ..Default::default()
381 };
382 assert_eq!(round_trip(&msg), msg);
383 }
384
385 #[test]
388 fn hop_zero_decodes_as_none() {
389 use crate::coding::Encode;
390
391 let version = Version::Lite05;
392 let mut params = Parameters::default();
393 params.set_varint(super::PARAM_HOP, 0);
394 let mut body = bytes::BytesMut::new();
395 params.encode(&mut body, version).unwrap();
396 let mut buf = bytes::BytesMut::new();
398 (body.len() as u64).encode(&mut buf, version).unwrap();
399 buf.extend_from_slice(&body);
400 let mut slice = &buf[..];
401 let got = Setup::decode(&mut slice, version).unwrap();
402 assert_eq!(got.hop, None);
403 }
404
405 #[test]
406 fn empty_path_round_trips() {
407 let msg = Setup {
410 path: Some(String::new()),
411 ..Default::default()
412 };
413 assert_eq!(round_trip(&msg), msg);
414 }
415
416 #[test]
417 fn roles_round_trip() {
418 for role in [Some(Role::Publisher), Some(Role::Subscriber), None] {
419 let msg = Setup {
420 path: Some("/room/123".to_string()),
421 role,
422 ..Default::default()
423 };
424 assert_eq!(round_trip(&msg), msg);
425 }
426 }
427
428 #[test]
429 fn unknown_probe_level_saturates_to_increase() {
430 let mut params = Parameters::default();
433 params.set_varint(PARAM_PROBE, 99);
434 let mut body = Vec::new();
435 params.encode(&mut body, Version::Lite05).unwrap();
436
437 let mut buf = bytes::BytesMut::new();
438 body.len().encode(&mut buf, Version::Lite05).unwrap();
439 buf.extend_from_slice(&body);
440
441 let mut slice = &buf[..];
442 let got = Setup::decode(&mut slice, Version::Lite05).unwrap();
443 assert_eq!(got.probe, ProbeLevel::Increase);
444 }
445
446 #[test]
447 fn role_wire_codes() {
448 for (role, code) in [(Role::Publisher, 1u64), (Role::Subscriber, 2)] {
451 assert_eq!(role.to_code(), code);
452 assert_eq!(Role::from_code(code), Some(role));
453 }
454 }
455
456 #[test]
457 fn unknown_role_decodes_as_bidirectional() {
458 for code in [0u64, 9, 250] {
462 let mut params = Parameters::default();
463 params.set_varint(PARAM_ROLE, code);
464 let mut body = Vec::new();
465 params.encode(&mut body, Version::Lite05).unwrap();
466
467 let mut buf = bytes::BytesMut::new();
468 body.len().encode(&mut buf, Version::Lite05).unwrap();
469 buf.extend_from_slice(&body);
470
471 let mut slice = &buf[..];
472 let got = Setup::decode(&mut slice, Version::Lite05).unwrap();
473 assert_eq!(got.role, None, "role code {code} should decode as bidirectional");
474 }
475 }
476
477 #[test]
478 fn rejects_before_lite05() {
479 let msg = Setup::default();
480 let mut buf = bytes::BytesMut::new();
481 assert!(matches!(
482 msg.encode(&mut buf, Version::Lite04),
483 Err(EncodeError::Version)
484 ));
485 }
486
487 #[test]
488 fn ignores_unknown_parameters() {
489 let mut params = Parameters::default();
491 params.set_bytes(PARAM_PATH, b"/foo".to_vec());
492 params.set_bytes(0xbeef, b"whatever".to_vec());
493
494 let mut body = Vec::new();
495 params.encode(&mut body, Version::Lite05).unwrap();
496
497 let mut buf = bytes::BytesMut::new();
499 body.len().encode(&mut buf, Version::Lite05).unwrap();
500 buf.extend_from_slice(&body);
501
502 let mut slice = &buf[..];
503 let got = Setup::decode(&mut slice, Version::Lite05).unwrap();
504 assert_eq!(got.path.as_deref(), Some("/foo"));
505 }
506}