1use crate::origin;
2#[cfg(test)]
3use crate::runtime::Timers;
4use crate::time::{Clock, Instant};
5use crate::{
6 ALPN_14, ALPN_15, ALPN_16, ALPN_17, ALPN_18, ALPN_19, ALPN_20, ALPN_21, ALPN_22, ALPN_LITE, ALPN_LITE_03,
7 ALPN_LITE_04, ALPN_LITE_05, ALPN_LITE_06, ALPN_LITE_07_WIP, Consume, Error, NEGOTIATED, Session, Version, Versions,
8 coding::{self, Decode, Encode, Stream},
9 ietf, lite, setup, stats,
10};
11
12#[derive(Default, Clone)]
14pub struct Client {
15 publish: Option<origin::Consumer>,
16 subscribe: Option<origin::Producer>,
17 stats: stats::Session,
18 versions: Versions,
19 setup_path: Option<String>,
20 setup_authority: Option<String>,
21 cost: Option<u64>,
22 peer_hop: Option<crate::Hop>,
23}
24
25impl Client {
26 pub fn new() -> Self {
28 Default::default()
29 }
30
31 pub fn with_publisher(mut self, publish: impl Consume<origin::Consumer>) -> Self {
35 self.publish = Some(publish.consume());
36 self
37 }
38
39 pub fn with_subscriber(mut self, subscribe: origin::Producer) -> Self {
42 self.subscribe = Some(subscribe);
43 self
44 }
45
46 pub fn with_stats(mut self, stats: stats::Session) -> Self {
51 self.stats = stats;
52 self
53 }
54
55 pub fn with_origin(self, origin: origin::Producer) -> Self {
60 self.with_publisher(&origin).with_subscriber(origin)
61 }
62
63 pub fn with_versions(mut self, versions: Versions) -> Self {
66 self.versions = versions;
67 self
68 }
69
70 pub fn with_path(mut self, path: impl Into<String>) -> Self {
82 self.setup_path = Some(path.into());
83 self
84 }
85
86 pub fn with_authority(mut self, authority: impl Into<String>) -> Self {
88 self.setup_authority = Some(authority.into());
89 self
90 }
91
92 pub fn with_cost(mut self, cost: u64) -> Self {
107 self.cost = Some(cost);
108 self
109 }
110
111 pub fn with_peer_hop(mut self, hop: crate::Hop) -> Self {
131 self.peer_hop = Some(hop);
132 self
133 }
134
135 fn origins(&self) -> (Option<origin::Consumer>, Option<origin::Producer>) {
146 if self.publish.is_none() && self.subscribe.is_none() {
147 tracing::warn!("not publishing or consuming anything");
148 }
149 let publish = self.publish.clone().map(|origin| origin.with_stats(self.stats.clone()));
150 let subscribe = self
151 .subscribe
152 .clone()
153 .map(|origin| origin.with_stats(self.stats.clone()));
154 let publish = publish.map(|origin| origin.excluding(self.peer_hop.unwrap_or(crate::Hop::UNKNOWN)));
155 (publish, subscribe)
156 }
157
158 fn start_lite<S>(
161 &self,
162 runtime: Clock,
163 session: S,
164 version: lite::Version,
165 ) -> Result<(Session, crate::Driver<S>), Error>
166 where
167 S: crate::transport::poll::Session,
168 {
169 let (publish, subscribe) = self.origins();
170
171 let our_setup = if version.has_setup_stream() {
177 lite::Setup {
178 probe: lite::ProbeLevel::detect(&session),
179 path: self.setup_path.clone(),
180 role: lite::Role::from_origins(self.publish.is_some(), self.subscribe.is_some()),
181 cost: self.cost,
182 hop: None,
184 }
185 } else {
186 lite::Setup::default()
187 };
188
189 let start = lite::start(lite::Config {
190 runtime: runtime.clone(),
191 session: session.clone(),
192 setup_stream: None,
193 publish,
194 subscribe,
195 peer_hop: self.peer_hop,
196 version,
197 our_setup,
198 peer_setup: None,
199 })?;
200
201 Ok(Session::new(
202 runtime,
203 session,
204 version.into(),
205 start.recv_bandwidth,
206 crate::driver::Protocol::Lite(Box::new(start.driver)),
207 start.goaway,
208 ))
209 }
210
211 pub async fn connect_lite<S>(&self, now: Instant, session: S) -> Result<(Session, crate::Driver<S>), Error>
221 where
222 S: crate::transport::poll::Session,
223 {
224 let runtime = Clock::new(now);
225 let version = match session.protocol() {
226 Some(ALPN_LITE_07_WIP) => lite::Version::Lite07,
227 Some(ALPN_LITE_06) => lite::Version::Lite06,
228 Some(ALPN_LITE_05) => lite::Version::Lite05,
229 Some(ALPN_LITE_04) => lite::Version::Lite04,
230 Some(ALPN_LITE_03) => lite::Version::Lite03,
231 _ => return Err(Error::Version),
232 };
233 self.versions.select(Version::Lite(version)).ok_or(Error::Version)?;
234 self.start_lite(runtime, session, version)
235 }
236
237 pub async fn connect<S>(&self, now: Instant, mut session: S) -> Result<(Session, crate::Driver<S>), Error>
241 where
242 S: crate::transport::poll::Boxable,
243 {
244 let runtime = Clock::new(now);
245 let (publish, subscribe) = self.origins();
246
247 let (encoding, supported) = match session.protocol() {
250 Some(alpn @ (ALPN_22 | ALPN_21 | ALPN_20 | ALPN_19 | ALPN_18 | ALPN_17)) => {
251 let draft = match alpn {
252 ALPN_22 => ietf::Version::Draft22,
253 ALPN_21 => ietf::Version::Draft21,
254 ALPN_20 => ietf::Version::Draft20,
255 ALPN_19 => ietf::Version::Draft19,
256 ALPN_18 => ietf::Version::Draft18,
257 _ => ietf::Version::Draft17,
258 };
259
260 let v = self.versions.select(Version::Ietf(draft)).ok_or(Error::Version)?;
261
262 let (protocol, goaway) = ietf::start(ietf::Config {
265 runtime: runtime.clone(),
266 session: session.clone(),
267 setup: None,
268 request_id_max: None,
269 client: true,
270 publish: publish.clone(),
271 subscribe: subscribe.clone(),
272 peer_hop: self.peer_hop,
273 cost: self.cost,
274 version: draft,
275 path: self.setup_path.clone(),
276 authority: self.setup_authority.clone(),
277 peer_setup_stream: None,
278 peer_declared: None,
279 })?;
280
281 tracing::debug!(version = ?v, "connected");
282 return Ok(Session::new(
283 runtime,
284 session,
285 v,
286 None,
287 crate::driver::Protocol::Ietf(protocol),
288 goaway,
289 ));
290 }
291 Some(ALPN_16) => {
292 let v = self
293 .versions
294 .select(Version::Ietf(ietf::Version::Draft16))
295 .ok_or(Error::Version)?;
296 (v, v.into())
297 }
298 Some(ALPN_15) => {
299 let v = self
300 .versions
301 .select(Version::Ietf(ietf::Version::Draft15))
302 .ok_or(Error::Version)?;
303 (v, v.into())
304 }
305 Some(ALPN_14) => {
306 let v = self
307 .versions
308 .select(Version::Ietf(ietf::Version::Draft14))
309 .ok_or(Error::Version)?;
310 (v, v.into())
311 }
312 Some(alpn @ (ALPN_LITE_05 | ALPN_LITE_06 | ALPN_LITE_07_WIP)) => {
313 let version = match alpn {
314 ALPN_LITE_07_WIP => lite::Version::Lite07,
315 ALPN_LITE_06 => lite::Version::Lite06,
316 _ => lite::Version::Lite05,
317 };
318 self.versions.select(Version::Lite(version)).ok_or(Error::Version)?;
319 return self.start_lite(runtime, session, version);
320 }
321 Some(ALPN_LITE_04) => {
322 self.versions
323 .select(Version::Lite(lite::Version::Lite04))
324 .ok_or(Error::Version)?;
325 return self.start_lite(runtime, session, lite::Version::Lite04);
326 }
327 Some(ALPN_LITE_03) => {
328 self.versions
329 .select(Version::Lite(lite::Version::Lite03))
330 .ok_or(Error::Version)?;
331 return self.start_lite(runtime, session, lite::Version::Lite03);
332 }
333 Some(ALPN_LITE) | None => {
334 let supported = self.versions.filter(&NEGOTIATED.into()).ok_or(Error::Version)?;
335 (Version::Ietf(ietf::Version::Draft14), supported)
336 }
337 Some(p) => return Err(Error::UnknownAlpn(p.to_string())),
338 };
339
340 let mut stream = Stream::open(&mut session, encoding).await?;
341
342 let ietf_encoding = ietf::Version::try_from(encoding).map_err(|_| Error::Version)?;
344
345 let mut parameters = ietf::Parameters::default();
346 parameters.set_varint(ietf::ParameterVarInt::MaxRequestId, u32::MAX as u64);
347 parameters.set_bytes(ietf::ParameterBytes::Implementation, b"moq-lite-rs".to_vec());
348 if let Some(path) = &self.setup_path {
350 parameters.set_bytes(ietf::ParameterBytes::Path, path.clone().into_bytes());
351 }
352 if let Some(authority) = &self.setup_authority {
353 parameters.set_bytes(ietf::ParameterBytes::Authority, authority.clone().into_bytes());
354 }
355 ietf::solicit::into_setup(&mut parameters, ietf_encoding);
356 ietf::hidden::into_setup(&mut parameters, ietf_encoding);
357 let parameters = parameters.encode_bytes(ietf_encoding)?;
358
359 let client = setup::Client {
360 versions: supported.clone().into(),
361 parameters,
362 };
363
364 stream.writer.encode(&client).await?;
365
366 let mut server: setup::Server = stream.reader.decode().await?;
367
368 let version = supported
369 .iter()
370 .find(|v| coding::Version::from(**v) == server.version)
371 .copied()
372 .ok_or(Error::Version)?;
373
374 let (recv_bw, protocol, goaway) = match version {
375 Version::Lite(v) => {
376 let stream = stream.with_version(v);
377 let start = lite::start(lite::Config {
378 runtime: runtime.clone(),
379 session: session.clone(),
380 setup_stream: Some(stream),
381 publish: publish.clone(),
382 subscribe: subscribe.clone(),
383 peer_hop: self.peer_hop,
384 version: v,
385 our_setup: lite::Setup::default(),
388 peer_setup: None,
389 })?;
390
391 (
392 start.recv_bandwidth,
393 crate::driver::Protocol::Lite(Box::new(start.driver)),
394 start.goaway,
395 )
396 }
397 Version::Ietf(v) => {
398 let parameters = ietf::Parameters::decode(&mut server.parameters, v)?;
401 let request_id_max = parameters
402 .get_varint(ietf::ParameterVarInt::MaxRequestId)
403 .map(ietf::RequestId);
404 let peer_declared = ietf::peer::Peer {
405 solicit: ietf::solicit::from_setup(¶meters, v)?,
406 hidden: ietf::hidden::from_setup(¶meters, v),
407 ..Default::default()
408 };
409
410 let stream = stream.with_version(v);
411 let (protocol, goaway) = ietf::start(ietf::Config {
413 runtime: runtime.clone(),
414 session: session.clone(),
415 setup: Some(stream),
416 request_id_max,
417 client: true,
418 publish: publish.clone(),
419 subscribe: subscribe.clone(),
420 peer_hop: self.peer_hop,
421 cost: self.cost,
422 version: v,
423 path: None,
424 authority: None,
425 peer_setup_stream: None,
426 peer_declared: Some(peer_declared),
427 })?;
428 (None, crate::driver::Protocol::Ietf(protocol), goaway)
429 }
430 };
431
432 Ok(Session::new(runtime, session, version, recv_bw, protocol, goaway))
433 }
434}
435
436#[cfg(test)]
437mod tests {
438 use super::*;
439 use crate::model::ProduceTest;
440 use std::{
441 collections::VecDeque,
442 sync::{Arc, Mutex},
443 };
444
445 use std::task::{Context, Poll};
446
447 use crate::SessionError;
448 use crate::coding::{Decode, Encode};
449 use bytes::{BufMut, Bytes};
450
451 #[derive(Debug, Clone, Default)]
452 struct FakeError;
453
454 impl std::fmt::Display for FakeError {
455 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
456 write!(f, "fake transport error")
457 }
458 }
459
460 impl std::error::Error for FakeError {}
461
462 impl web_transport_trait::Error for FakeError {
463 fn session_error(&self) -> Option<(u32, String)> {
464 Some((0, "closed".to_string()))
465 }
466 }
467
468 #[derive(Clone, Default)]
469 struct FakeSession {
470 state: Arc<FakeSessionState>,
471 park: kio::Park,
473 }
474
475 #[derive(Default)]
476 struct FakeSessionState {
477 protocol: Option<&'static str>,
478 control_stream: Mutex<Option<(FakeSendStream, FakeRecvStream)>>,
479 close_events: Mutex<Vec<(u32, String)>>,
480 closed: kio::Fan,
481 control_writes: Arc<Mutex<Vec<u8>>>,
482 send_rate: Mutex<Option<u64>>,
483 bytes_sent: Mutex<Option<u64>>,
484 }
485
486 impl FakeSession {
487 fn new(protocol: Option<&'static str>, server_control_bytes: Vec<u8>) -> Self {
488 let writes = Arc::new(Mutex::new(Vec::new()));
489 let send = FakeSendStream { writes: writes.clone() };
490 let recv = FakeRecvStream {
491 data: VecDeque::from(server_control_bytes),
492 };
493 let state = FakeSessionState {
494 protocol,
495 control_stream: Mutex::new(Some((send, recv))),
496 close_events: Mutex::new(Vec::new()),
497 closed: kio::Fan::default(),
498 control_writes: writes,
499 send_rate: Mutex::new(None),
500 bytes_sent: Mutex::new(None),
501 };
502 Self {
503 state: Arc::new(state),
504 park: kio::Park::default(),
505 }
506 }
507
508 fn set_send_rate(&self, rate: Option<u64>) {
509 *self.state.send_rate.lock().unwrap() = rate;
510 }
511
512 fn set_bytes_sent(&self, bytes: Option<u64>) {
513 *self.state.bytes_sent.lock().unwrap() = bytes;
514 }
515
516 fn control_writes(&self) -> Vec<u8> {
517 self.state.control_writes.lock().unwrap().clone()
518 }
519
520 async fn wait_for_first_close(&self) -> (u32, String) {
521 kio::wait(|waiter| {
522 self.state.closed.register(waiter);
523 match self.state.close_events.lock().unwrap().first().cloned() {
524 Some(close) => std::task::Poll::Ready(close),
525 None => std::task::Poll::Pending,
526 }
527 })
528 .await
529 }
530 }
531
532 impl web_transport_trait::poll::Session for FakeSession {
533 type SendStream = FakeSendStream;
534 type RecvStream = FakeRecvStream;
535 type Error = FakeError;
536
537 fn poll_accept_uni(&mut self, _cx: &mut Context<'_>) -> Poll<Result<Self::RecvStream, Self::Error>> {
538 Poll::Pending
539 }
540
541 fn poll_accept_bi(
542 &mut self,
543 _cx: &mut Context<'_>,
544 ) -> Poll<Result<(Self::SendStream, Self::RecvStream), Self::Error>> {
545 Poll::Pending
546 }
547
548 fn poll_open_bi(
549 &mut self,
550 _cx: &mut Context<'_>,
551 ) -> Poll<Result<(Self::SendStream, Self::RecvStream), Self::Error>> {
552 Poll::Ready(self.state.control_stream.lock().unwrap().take().ok_or(FakeError))
553 }
554
555 fn poll_open_uni(&mut self, _cx: &mut Context<'_>) -> Poll<Result<Self::SendStream, Self::Error>> {
556 Poll::Pending
557 }
558
559 fn poll_send_datagram(&mut self, _cx: &mut Context<'_>, _payload: &[u8]) -> Poll<Result<(), Self::Error>> {
560 Poll::Ready(Ok(()))
561 }
562
563 fn poll_recv_datagram(&mut self, _cx: &mut Context<'_>) -> Poll<Result<Bytes, Self::Error>> {
564 Poll::Pending
565 }
566
567 fn max_datagram_size(&self) -> usize {
568 1200
569 }
570
571 fn protocol(&self) -> Option<&str> {
572 self.state.protocol
573 }
574
575 fn close(&mut self, code: u32, reason: &str) {
576 self.state.close_events.lock().unwrap().push((code, reason.to_string()));
577 self.state.closed.wake();
578 }
579
580 fn poll_closed(&mut self, cx: &mut Context<'_>) -> Poll<Self::Error> {
581 self.state.closed.register(self.park.hold(cx));
583 match self.state.close_events.lock().unwrap().is_empty() {
584 false => Poll::Ready(FakeError),
585 true => Poll::Pending,
586 }
587 }
588
589 fn stats(&self) -> impl web_transport_trait::Stats {
590 FakeStats {
591 send_rate: *self.state.send_rate.lock().unwrap(),
592 bytes_sent: *self.state.bytes_sent.lock().unwrap(),
593 }
594 }
595 }
596
597 struct FakeStats {
598 send_rate: Option<u64>,
599 bytes_sent: Option<u64>,
600 }
601
602 impl web_transport_trait::Stats for FakeStats {
603 fn estimated_send_rate(&self) -> Option<u64> {
604 self.send_rate
605 }
606
607 fn bytes_sent(&self) -> Option<u64> {
608 self.bytes_sent
609 }
610 }
611
612 #[derive(Clone, Default)]
613 struct FakeSendStream {
614 writes: Arc<Mutex<Vec<u8>>>,
615 }
616
617 impl web_transport_trait::poll::SendStream for FakeSendStream {
618 type Error = FakeError;
619
620 fn poll_write(&mut self, _cx: &mut Context<'_>, buf: &[u8]) -> Poll<Result<usize, Self::Error>> {
621 self.writes.lock().unwrap().put_slice(buf);
622 Poll::Ready(Ok(buf.len()))
623 }
624
625 fn set_priority(&mut self, _order: u8) {}
626
627 fn finish(&mut self) -> Result<(), Self::Error> {
628 Ok(())
629 }
630
631 fn reset(&mut self, _code: u32) {}
632
633 fn poll_closed(&mut self, _cx: &mut Context<'_>) -> Poll<Result<(), Self::Error>> {
634 Poll::Ready(Ok(()))
635 }
636 }
637
638 struct FakeRecvStream {
639 data: VecDeque<u8>,
640 }
641
642 impl web_transport_trait::poll::RecvStream for FakeRecvStream {
643 type Error = FakeError;
644
645 fn poll_read(&mut self, _cx: &mut Context<'_>, dst: &mut [u8]) -> Poll<Result<Option<usize>, Self::Error>> {
646 if self.data.is_empty() {
647 return Poll::Ready(Ok(None));
648 }
649
650 let size = dst.len().min(self.data.len());
651 for slot in dst.iter_mut().take(size) {
652 *slot = self.data.pop_front().unwrap();
653 }
654 Poll::Ready(Ok(Some(size)))
655 }
656
657 fn stop(&mut self, _code: u32) {}
658
659 fn poll_closed(&mut self, _cx: &mut Context<'_>) -> Poll<Result<(), Self::Error>> {
660 Poll::Ready(Ok(()))
661 }
662 }
663
664 fn mock_server_setup(negotiated: Version) -> Vec<u8> {
665 let mut encoded = Vec::new();
666 let server = setup::Server {
667 version: negotiated.into(),
668 parameters: Bytes::new(),
669 };
670 server
671 .encode(&mut encoded, Version::Ietf(ietf::Version::Draft14))
672 .unwrap();
673
674 let info = lite::SessionInfo { bitrate: Some(1) };
676 let lite_v = lite::Version::try_from(negotiated).unwrap();
677 info.encode(&mut encoded, lite_v).unwrap();
678
679 encoded
680 }
681
682 async fn run_alpn_lite_fallback_case(protocol: Option<&'static str>) {
683 let fake = FakeSession::new(protocol, mock_server_setup(Version::Lite(lite::Version::Lite01)));
684 let client = Client::new().with_versions(
685 [
686 Version::Lite(lite::Version::Lite03),
687 Version::Lite(lite::Version::Lite02),
688 Version::Lite(lite::Version::Lite01),
689 Version::Ietf(ietf::Version::Draft14),
690 ]
691 .into(),
692 );
693
694 let (_session, driver) = client
696 .connect(tokio::time::Instant::now().into_std(), fake.clone())
697 .await
698 .unwrap();
699 tokio::spawn(crate::time::run(driver));
700
701 let mut setup_bytes = Bytes::from(fake.control_writes());
703 let setup = setup::Client::decode(&mut setup_bytes, Version::Ietf(ietf::Version::Draft14)).unwrap();
704 let advertised: Vec<Version> = setup.versions.iter().map(|v| Version::try_from(*v).unwrap()).collect();
705 assert_eq!(
706 advertised,
707 vec![
708 Version::Lite(lite::Version::Lite02),
709 Version::Lite(lite::Version::Lite01),
710 Version::Ietf(ietf::Version::Draft14),
711 ]
712 );
713
714 let (code, _) = fake.wait_for_first_close().await;
721 assert_ne!(code, SessionError::Version.to_code(), "SessionInfo failed to decode");
723 }
724
725 #[tokio::test(start_paused = true)]
731 async fn connect_does_not_wait_for_the_peer_to_announce() {
732 let gate = kio::Producer::new(true);
734 let transport = crate::lite::test_transport::SinkSession::gated_bi(gate.consume())
735 .with_protocol(crate::version::ALPN_LITE_05);
736
737 let origin = crate::origin::Config::new(crate::Hop::new(1).unwrap()).produce();
739 let client = Client::new()
740 .with_versions([Version::Lite(lite::Version::Lite05)].into())
741 .with_subscriber(origin);
742
743 let (_session, _driver) = tokio::time::timeout(
746 std::time::Duration::from_secs(30),
747 client.connect(tokio::time::Instant::now().into_std(), transport),
748 )
749 .await
750 .expect("connect waited on a peer that never announced")
751 .expect("connect failed");
752 }
753
754 #[tokio::test(start_paused = true)]
757 async fn draft14_setup_carries_the_authority() {
758 let fake = FakeSession::new(Some(ALPN_LITE), mock_server_setup(Version::Lite(lite::Version::Lite01)));
759 let client = Client::new()
760 .with_versions(
761 [
762 Version::Lite(lite::Version::Lite01),
763 Version::Ietf(ietf::Version::Draft14),
764 ]
765 .into(),
766 )
767 .with_path("/anon")
768 .with_authority("relay.example.com:4443");
769
770 let (_session, driver) = client
771 .connect(tokio::time::Instant::now().into_std(), fake.clone())
772 .await
773 .unwrap();
774 tokio::spawn(crate::time::run(driver));
775
776 let mut setup_bytes = Bytes::from(fake.control_writes());
777 let setup = setup::Client::decode(&mut setup_bytes, Version::Ietf(ietf::Version::Draft14)).unwrap();
778 let mut parameters = setup.parameters;
779 let parameters = ietf::Parameters::decode(&mut parameters, ietf::Version::Draft14).unwrap();
780 assert_eq!(
781 parameters.get_bytes(ietf::ParameterBytes::Authority),
782 Some(b"relay.example.com:4443".as_ref())
783 );
784 assert_eq!(
785 parameters.get_bytes(ietf::ParameterBytes::Path),
786 Some(b"/anon".as_ref())
787 );
788 }
789
790 #[tokio::test(start_paused = true)]
791 async fn alpn_lite_falls_back_to_draft14_and_switches_version_post_setup() {
792 run_alpn_lite_fallback_case(Some(ALPN_LITE)).await;
793 }
794
795 #[tokio::test(start_paused = true)]
796 async fn no_alpn_falls_back_to_draft14_and_switches_version_post_setup() {
797 run_alpn_lite_fallback_case(None).await;
798 }
799
800 #[test]
803 fn driver_is_caller_polled_and_holds_no_session() {
804 let fake = FakeSession::new(Some(ALPN_LITE_04), Vec::new());
805 let client = Client::new().with_versions(Version::Lite(lite::Version::Lite04).into());
806
807 let runtime = crate::runtime::Test::new();
808 let (session, mut driver) = futures::executor::block_on(client.connect(runtime.now(), fake.clone())).unwrap();
809 assert_eq!(session.version(), Version::Lite(lite::Version::Lite04));
810
811 assert!(driver.poll(runtime.now(), &kio::Waiter::noop()).is_ok());
813
814 drop(session);
817 assert!(fake.state.close_events.lock().unwrap().is_empty());
818 let _ = driver.poll(runtime.now(), &kio::Waiter::noop());
819 assert_eq!(
820 fake.state.close_events.lock().unwrap()[0].0,
821 SessionError::Cancel.to_code()
822 );
823 }
824
825 #[test]
829 fn session_clones_share_the_close() {
830 let fake = FakeSession::new(Some(ALPN_LITE_04), Vec::new());
831 let client = Client::new().with_versions(Version::Lite(lite::Version::Lite04).into());
832
833 let runtime = crate::runtime::Test::new();
834 let (session, mut driver) = futures::executor::block_on(client.connect(runtime.now(), fake.clone())).unwrap();
835 let clone = session.clone();
836
837 drop(session);
839 assert!(fake.state.close_events.lock().unwrap().is_empty());
840 let _ = driver.poll(runtime.now(), &kio::Waiter::noop());
841 assert!(fake.state.close_events.lock().unwrap().is_empty());
842
843 clone.abort(Error::Cancel);
844 let _ = driver.poll(runtime.now(), &kio::Waiter::noop());
845 assert_eq!(
846 fake.state.close_events.lock().unwrap()[0].0,
847 SessionError::Cancel.to_code()
848 );
849
850 let _ = driver.poll(runtime.now(), &kio::Waiter::noop());
853 futures::executor::block_on(clone.closed());
854
855 let closes = fake.state.close_events.lock().unwrap().len();
857 drop(clone);
858 let _ = driver.poll(runtime.now(), &kio::Waiter::noop());
859 assert_eq!(fake.state.close_events.lock().unwrap().len(), closes);
860 }
861
862 #[test]
866 fn dropped_driver_resolves_closed() {
867 let fake = FakeSession::new(Some(ALPN_LITE_04), Vec::new());
868 let client = Client::new().with_versions(Version::Lite(lite::Version::Lite04).into());
869
870 let runtime = crate::runtime::Test::new();
871 let (session, driver) = futures::executor::block_on(client.connect(runtime.now(), fake.clone())).unwrap();
872
873 drop(driver);
874 assert!(matches!(futures::executor::block_on(session.closed()), Error::Cancel));
875 }
876
877 #[derive(Clone)]
881 struct LocalSession {
882 inner: FakeSession,
883 _local: std::rc::Rc<()>,
884 }
885
886 struct LocalSend {
887 inner: FakeSendStream,
888 _local: std::rc::Rc<()>,
889 }
890
891 struct LocalRecv {
892 inner: FakeRecvStream,
893 _local: std::rc::Rc<()>,
894 }
895
896 impl web_transport_trait::poll::Session for LocalSession {
897 type SendStream = LocalSend;
898 type RecvStream = LocalRecv;
899 type Error = FakeError;
900
901 fn poll_accept_uni(&mut self, cx: &mut Context<'_>) -> Poll<Result<Self::RecvStream, Self::Error>> {
902 self.inner.poll_accept_uni(cx).map_ok(|stream| LocalRecv {
903 inner: stream,
904 _local: self._local.clone(),
905 })
906 }
907
908 fn poll_accept_bi(
909 &mut self,
910 cx: &mut Context<'_>,
911 ) -> Poll<Result<(Self::SendStream, Self::RecvStream), Self::Error>> {
912 self.inner.poll_accept_bi(cx).map_ok(|(send, recv)| {
913 (
914 LocalSend {
915 inner: send,
916 _local: self._local.clone(),
917 },
918 LocalRecv {
919 inner: recv,
920 _local: self._local.clone(),
921 },
922 )
923 })
924 }
925
926 fn poll_open_bi(
927 &mut self,
928 cx: &mut Context<'_>,
929 ) -> Poll<Result<(Self::SendStream, Self::RecvStream), Self::Error>> {
930 self.inner.poll_open_bi(cx).map_ok(|(send, recv)| {
931 (
932 LocalSend {
933 inner: send,
934 _local: self._local.clone(),
935 },
936 LocalRecv {
937 inner: recv,
938 _local: self._local.clone(),
939 },
940 )
941 })
942 }
943
944 fn poll_open_uni(&mut self, cx: &mut Context<'_>) -> Poll<Result<Self::SendStream, Self::Error>> {
945 self.inner.poll_open_uni(cx).map_ok(|stream| LocalSend {
946 inner: stream,
947 _local: self._local.clone(),
948 })
949 }
950
951 fn poll_send_datagram(&mut self, cx: &mut Context<'_>, payload: &[u8]) -> Poll<Result<(), Self::Error>> {
952 self.inner.poll_send_datagram(cx, payload)
953 }
954
955 fn poll_recv_datagram(&mut self, cx: &mut Context<'_>) -> Poll<Result<Bytes, Self::Error>> {
956 self.inner.poll_recv_datagram(cx)
957 }
958
959 fn max_datagram_size(&self) -> usize {
960 self.inner.max_datagram_size()
961 }
962
963 fn protocol(&self) -> Option<&str> {
964 self.inner.protocol()
965 }
966
967 fn close(&mut self, code: u32, reason: &str) {
968 self.inner.close(code, reason);
969 }
970
971 fn poll_closed(&mut self, cx: &mut Context<'_>) -> Poll<Self::Error> {
972 self.inner.poll_closed(cx)
973 }
974
975 fn stats(&self) -> impl web_transport_trait::Stats {
976 self.inner.stats()
977 }
978 }
979
980 impl web_transport_trait::poll::SendStream for LocalSend {
981 type Error = FakeError;
982
983 fn poll_write(&mut self, cx: &mut Context<'_>, buf: &[u8]) -> Poll<Result<usize, Self::Error>> {
984 self.inner.poll_write(cx, buf)
985 }
986
987 fn set_priority(&mut self, order: u8) {
988 self.inner.set_priority(order);
989 }
990
991 fn finish(&mut self) -> Result<(), Self::Error> {
992 self.inner.finish()
993 }
994
995 fn reset(&mut self, code: u32) {
996 self.inner.reset(code);
997 }
998
999 fn poll_closed(&mut self, cx: &mut Context<'_>) -> Poll<Result<(), Self::Error>> {
1000 web_transport_trait::poll::SendStream::poll_closed(&mut self.inner, cx)
1001 }
1002 }
1003
1004 impl web_transport_trait::poll::RecvStream for LocalRecv {
1005 type Error = FakeError;
1006
1007 fn poll_read(&mut self, cx: &mut Context<'_>, dst: &mut [u8]) -> Poll<Result<Option<usize>, Self::Error>> {
1008 self.inner.poll_read(cx, dst)
1009 }
1010
1011 fn stop(&mut self, code: u32) {
1012 self.inner.stop(code);
1013 }
1014
1015 fn poll_closed(&mut self, cx: &mut Context<'_>) -> Poll<Result<(), Self::Error>> {
1016 web_transport_trait::poll::RecvStream::poll_closed(&mut self.inner, cx)
1017 }
1018 }
1019
1020 #[test]
1025 fn connect_lite_over_a_send_less_transport() {
1026 let fake = FakeSession::new(Some(ALPN_LITE_04), Vec::new());
1027 let local = LocalSession {
1028 inner: fake.clone(),
1029 _local: std::rc::Rc::new(()),
1030 };
1031 let client = Client::new().with_versions(Version::Lite(lite::Version::Lite04).into());
1032
1033 let runtime = crate::runtime::Test::new();
1034 let (session, mut driver) = futures::executor::block_on(client.connect_lite(runtime.now(), local)).unwrap();
1035 assert!(driver.poll(runtime.now(), &kio::Waiter::noop()).is_ok());
1036
1037 fn assert_send_sync<T: Send + Sync>(_: &T) {}
1038 assert_send_sync(&session);
1039
1040 session.abort(Error::Cancel);
1041 let _ = driver.poll(runtime.now(), &kio::Waiter::noop());
1042 assert_eq!(
1043 fake.state.close_events.lock().unwrap()[0].0,
1044 SessionError::Cancel.to_code()
1045 );
1046 }
1047
1048 #[test]
1051 fn accept_lite_over_a_send_less_transport() {
1052 let fake = FakeSession::new(Some(ALPN_LITE_04), Vec::new());
1053 let local = LocalSession {
1054 inner: fake.clone(),
1055 _local: std::rc::Rc::new(()),
1056 };
1057 let server = crate::Server::new().with_versions(Version::Lite(lite::Version::Lite04).into());
1058
1059 let runtime = crate::runtime::Test::new();
1060 let (session, mut driver) = futures::executor::block_on(server.accept_lite(runtime.now(), local)).unwrap();
1061 assert_eq!(session.version(), Version::Lite(lite::Version::Lite04));
1062 assert!(driver.poll(runtime.now(), &kio::Waiter::noop()).is_ok());
1063
1064 drop(session);
1065 assert!(fake.state.close_events.lock().unwrap().is_empty());
1066 let _ = driver.poll(runtime.now(), &kio::Waiter::noop());
1067 assert_eq!(
1068 fake.state.close_events.lock().unwrap()[0].0,
1069 SessionError::Cancel.to_code()
1070 );
1071 }
1072
1073 #[test]
1076 fn connect_lite_refuses_ietf_alpns() {
1077 let fake = FakeSession::new(Some(ALPN_19), Vec::new());
1078 let local = LocalSession {
1079 inner: fake,
1080 _local: std::rc::Rc::new(()),
1081 };
1082 let client = Client::new();
1083 let runtime = crate::runtime::Test::new();
1084 let result = futures::executor::block_on(client.connect_lite(runtime.now(), local));
1085 assert!(matches!(result, Err(Error::Version)));
1086 }
1087
1088 #[tokio::test(start_paused = true)]
1092 async fn stats_reads_prime_the_sampler() {
1093 let fake = FakeSession::new(Some(ALPN_LITE_04), Vec::new());
1094 fake.set_send_rate(Some(1_000_000));
1095
1096 let client = Client::new().with_versions(Version::Lite(lite::Version::Lite04).into());
1097 let (session, driver) = client
1098 .connect(tokio::time::Instant::now().into_std(), fake.clone())
1099 .await
1100 .unwrap();
1101 tokio::spawn(crate::time::run(driver));
1102
1103 assert_eq!(
1105 session.stats().estimated_send_rate,
1106 Some(crate::bandwidth::Rate::from_bps(1_000_000))
1107 );
1108
1109 fake.set_send_rate(Some(2_000_000));
1112 while session.stats().estimated_send_rate != Some(crate::bandwidth::Rate::from_bps(2_000_000)) {
1113 tokio::time::sleep(std::time::Duration::from_millis(10)).await;
1114 }
1115 }
1116
1117 #[tokio::test(start_paused = true)]
1124 async fn stats_capture_the_final_counters() {
1125 let fake = FakeSession::new(Some(ALPN_LITE_04), Vec::new());
1126 fake.set_send_rate(None);
1127 fake.set_bytes_sent(Some(0));
1128
1129 let client = Client::new().with_versions(Version::Lite(lite::Version::Lite04).into());
1130 let (session, driver) = client
1131 .connect(tokio::time::Instant::now().into_std(), fake.clone())
1132 .await
1133 .unwrap();
1134 tokio::spawn(crate::time::run(driver));
1135 assert!(
1136 session.send_bandwidth().is_none(),
1137 "no send-rate estimate, so nothing samples on its own"
1138 );
1139
1140 fake.set_bytes_sent(Some(4242));
1141
1142 session.abort(Error::Cancel);
1143 session.closed().await;
1144
1145 assert_eq!(
1146 session.stats().bytes_sent,
1147 Some(4242),
1148 "the closing snapshot must carry the session's final counters"
1149 );
1150 }
1151
1152 #[tokio::test(start_paused = true)]
1156 async fn send_bandwidth_samples_while_the_driver_runs() {
1157 let fake = FakeSession::new(Some(ALPN_LITE_04), Vec::new());
1158 fake.set_send_rate(Some(1_000_000));
1159
1160 let client = Client::new().with_versions(Version::Lite(lite::Version::Lite04).into());
1161 let (session, driver) = client
1162 .connect(tokio::time::Instant::now().into_std(), fake.clone())
1163 .await
1164 .unwrap();
1165 tokio::spawn(crate::time::run(driver));
1166
1167 let mut bandwidth = session.send_bandwidth().expect("backend reports an estimate");
1168 assert_eq!(
1169 bandwidth.changed().await.unwrap(),
1170 Some(crate::bandwidth::Rate::from_bps(1_000_000))
1171 );
1172
1173 fake.set_send_rate(Some(2_000_000));
1175 assert_eq!(
1176 bandwidth.changed().await.unwrap(),
1177 Some(crate::bandwidth::Rate::from_bps(2_000_000))
1178 );
1179 }
1180}