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,
41 Increase,
43}
44
45impl ProbeLevel {
46 pub fn detect<S: web_transport_trait::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 origin: Option<crate::Origin>,
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 origin = params
211 .get_varint(PARAM_ORIGIN)?
212 .and_then(|id| crate::Origin::new(id).ok());
213
214 Ok(Self {
215 probe,
216 path,
217 role,
218 cost,
219 origin,
220 })
221 }
222
223 fn encode_msg<W: bytes::BufMut>(&self, w: &mut W, version: Version) -> Result<(), EncodeError> {
224 if !version.has_setup_stream() {
225 return Err(EncodeError::Version);
226 }
227
228 let mut params = Parameters::default();
229 if self.probe != ProbeLevel::None {
231 params.set_varint(PARAM_PROBE, self.probe.to_code());
232 }
233 if let Some(path) = &self.path {
234 params.set_bytes(PARAM_PATH, path.as_bytes().to_vec());
235 }
236 if let Some(role) = self.role {
239 params.set_varint(PARAM_ROLE, role.to_code());
240 }
241 if let Some(cost) = self.cost {
242 params.set_varint(PARAM_COST, cost);
243 }
244 if let Some(origin) = self.origin {
245 params.set_varint(PARAM_ORIGIN, origin.id());
246 }
247
248 params.encode(w, version)
249 }
250}
251
252#[derive(Clone, Default)]
258pub(crate) struct PeerSetup(kio::Shared<Option<Setup>>);
259
260impl PeerSetup {
261 pub fn set(&self, setup: Setup) {
263 *self.0.lock() = Some(setup);
264 }
265
266 pub async fn probe_level(&self) -> ProbeLevel {
268 self.wait(|setup| setup.probe).await
269 }
270
271 pub async fn cost(&self) -> Option<u64> {
274 self.wait(|setup| setup.cost).await
275 }
276
277 pub async fn origin(&self) -> Option<crate::Origin> {
280 self.wait(|setup| setup.origin).await
281 }
282
283 async fn wait<T>(&self, f: impl FnOnce(&Setup) -> T) -> T {
289 let slot = self
290 .0
291 .wait(|setup| {
292 if setup.is_some() {
293 std::task::Poll::Ready(())
294 } else {
295 std::task::Poll::Pending
296 }
297 })
298 .await;
299 f(slot.as_ref().expect("waited for Some"))
300 }
301}
302
303#[cfg(test)]
304mod tests {
305 use super::*;
306
307 fn round_trip(msg: &Setup) -> Setup {
308 let mut buf = bytes::BytesMut::new();
309 msg.encode(&mut buf, Version::Lite05).unwrap();
310 let mut slice = &buf[..];
311 let got = Setup::decode(&mut slice, Version::Lite05).unwrap();
312 assert!(bytes::Buf::remaining(&slice) == 0, "trailing bytes after decode");
313 got
314 }
315
316 #[test]
317 fn empty_round_trip() {
318 let msg = Setup::default();
319 assert_eq!(round_trip(&msg), msg);
320 }
321
322 #[test]
326 fn detect_reports_nothing_without_stats() {
327 use crate::lite::test_transport::{SinkSession, SinkStats};
328 let session = SinkSession::new(Default::default()).with_stats(SinkStats::default());
329 assert_eq!(ProbeLevel::detect(&session), ProbeLevel::None);
330 }
331
332 #[test]
335 fn detect_reports_with_either_metric() {
336 use crate::lite::test_transport::{SinkSession, SinkStats};
337
338 let rtt_only = SinkStats::default().with_rtt(std::time::Duration::from_millis(40));
339 let session = SinkSession::new(Default::default()).with_stats(rtt_only);
340 assert_eq!(ProbeLevel::detect(&session), ProbeLevel::Report);
341
342 let rate_only = SinkStats::default().with_send_rate(1_000_000);
343 let session = SinkSession::new(Default::default()).with_stats(rate_only);
344 assert_eq!(ProbeLevel::detect(&session), ProbeLevel::Report);
345 }
346
347 #[test]
348 fn probe_levels_round_trip() {
349 for probe in [ProbeLevel::None, ProbeLevel::Report, ProbeLevel::Increase] {
350 let msg = Setup {
351 probe,
352 ..Default::default()
353 };
354 assert_eq!(round_trip(&msg), msg);
355 }
356 }
357
358 #[test]
359 fn cost_round_trip() {
360 for cost in [None, Some(0), Some(1), Some(7)] {
363 let msg = Setup {
364 cost,
365 ..Default::default()
366 };
367 assert_eq!(round_trip(&msg), msg);
368 }
369 }
370
371 #[test]
372 fn path_round_trip() {
373 let msg = Setup {
374 probe: ProbeLevel::Report,
375 path: Some("/room/123".to_string()),
376 ..Default::default()
377 };
378 assert_eq!(round_trip(&msg), msg);
379 }
380
381 #[test]
382 fn origin_round_trip() {
383 let msg = Setup {
384 origin: Some(crate::Origin::new(42).unwrap()),
385 ..Default::default()
386 };
387 assert_eq!(round_trip(&msg), msg);
388 }
389
390 #[test]
393 fn origin_zero_decodes_as_none() {
394 use crate::coding::Encode;
395
396 let version = Version::Lite05;
397 let mut params = Parameters::default();
398 params.set_varint(super::PARAM_ORIGIN, 0);
399 let mut body = bytes::BytesMut::new();
400 params.encode(&mut body, version).unwrap();
401 let mut buf = bytes::BytesMut::new();
403 (body.len() as u64).encode(&mut buf, version).unwrap();
404 buf.extend_from_slice(&body);
405 let mut slice = &buf[..];
406 let got = Setup::decode(&mut slice, version).unwrap();
407 assert_eq!(got.origin, None);
408 }
409
410 #[test]
411 fn empty_path_round_trips() {
412 let msg = Setup {
415 path: Some(String::new()),
416 ..Default::default()
417 };
418 assert_eq!(round_trip(&msg), msg);
419 }
420
421 #[test]
422 fn roles_round_trip() {
423 for role in [Some(Role::Publisher), Some(Role::Subscriber), None] {
424 let msg = Setup {
425 path: Some("/room/123".to_string()),
426 role,
427 ..Default::default()
428 };
429 assert_eq!(round_trip(&msg), msg);
430 }
431 }
432
433 #[test]
434 fn unknown_probe_level_saturates_to_increase() {
435 let mut params = Parameters::default();
438 params.set_varint(PARAM_PROBE, 99);
439 let mut body = Vec::new();
440 params.encode(&mut body, Version::Lite05).unwrap();
441
442 let mut buf = bytes::BytesMut::new();
443 body.len().encode(&mut buf, Version::Lite05).unwrap();
444 buf.extend_from_slice(&body);
445
446 let mut slice = &buf[..];
447 let got = Setup::decode(&mut slice, Version::Lite05).unwrap();
448 assert_eq!(got.probe, ProbeLevel::Increase);
449 }
450
451 #[test]
452 fn role_wire_codes() {
453 for (role, code) in [(Role::Publisher, 1u64), (Role::Subscriber, 2)] {
456 assert_eq!(role.to_code(), code);
457 assert_eq!(Role::from_code(code), Some(role));
458 }
459 }
460
461 #[test]
462 fn unknown_role_decodes_as_bidirectional() {
463 for code in [0u64, 9, 250] {
467 let mut params = Parameters::default();
468 params.set_varint(PARAM_ROLE, code);
469 let mut body = Vec::new();
470 params.encode(&mut body, Version::Lite05).unwrap();
471
472 let mut buf = bytes::BytesMut::new();
473 body.len().encode(&mut buf, Version::Lite05).unwrap();
474 buf.extend_from_slice(&body);
475
476 let mut slice = &buf[..];
477 let got = Setup::decode(&mut slice, Version::Lite05).unwrap();
478 assert_eq!(got.role, None, "role code {code} should decode as bidirectional");
479 }
480 }
481
482 #[test]
483 fn rejects_before_lite05() {
484 let msg = Setup::default();
485 let mut buf = bytes::BytesMut::new();
486 assert!(matches!(
487 msg.encode(&mut buf, Version::Lite04),
488 Err(EncodeError::Version)
489 ));
490 }
491
492 #[test]
493 fn ignores_unknown_parameters() {
494 let mut params = Parameters::default();
496 params.set_bytes(PARAM_PATH, b"/foo".to_vec());
497 params.set_bytes(0xbeef, b"whatever".to_vec());
498
499 let mut body = Vec::new();
500 params.encode(&mut body, Version::Lite05).unwrap();
501
502 let mut buf = bytes::BytesMut::new();
504 body.len().encode(&mut buf, Version::Lite05).unwrap();
505 buf.extend_from_slice(&body);
506
507 let mut slice = &buf[..];
508 let got = Setup::decode(&mut slice, Version::Lite05).unwrap();
509 assert_eq!(got.path.as_deref(), Some("/foo"));
510 }
511}