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 early_unis: Vec::new(),
280 })?;
281
282 tracing::debug!(version = ?v, "connected");
283 return Ok(Session::new(
284 runtime,
285 session,
286 v,
287 None,
288 crate::driver::Protocol::Ietf(protocol),
289 goaway,
290 ));
291 }
292 Some(ALPN_16) => {
293 let v = self
294 .versions
295 .select(Version::Ietf(ietf::Version::Draft16))
296 .ok_or(Error::Version)?;
297 (v, v.into())
298 }
299 Some(ALPN_15) => {
300 let v = self
301 .versions
302 .select(Version::Ietf(ietf::Version::Draft15))
303 .ok_or(Error::Version)?;
304 (v, v.into())
305 }
306 Some(ALPN_14) => {
307 let v = self
308 .versions
309 .select(Version::Ietf(ietf::Version::Draft14))
310 .ok_or(Error::Version)?;
311 (v, v.into())
312 }
313 Some(alpn @ (ALPN_LITE_05 | ALPN_LITE_06 | ALPN_LITE_07_WIP)) => {
314 let version = match alpn {
315 ALPN_LITE_07_WIP => lite::Version::Lite07,
316 ALPN_LITE_06 => lite::Version::Lite06,
317 _ => lite::Version::Lite05,
318 };
319 self.versions.select(Version::Lite(version)).ok_or(Error::Version)?;
320 return self.start_lite(runtime, session, version);
321 }
322 Some(ALPN_LITE_04) => {
323 self.versions
324 .select(Version::Lite(lite::Version::Lite04))
325 .ok_or(Error::Version)?;
326 return self.start_lite(runtime, session, lite::Version::Lite04);
327 }
328 Some(ALPN_LITE_03) => {
329 self.versions
330 .select(Version::Lite(lite::Version::Lite03))
331 .ok_or(Error::Version)?;
332 return self.start_lite(runtime, session, lite::Version::Lite03);
333 }
334 Some(ALPN_LITE) | None => {
335 let supported = self.versions.filter(&NEGOTIATED.into()).ok_or(Error::Version)?;
336 (Version::Ietf(ietf::Version::Draft14), supported)
337 }
338 Some(p) => return Err(Error::UnknownAlpn(p.to_string())),
339 };
340
341 let mut stream = Stream::open(&mut session, encoding).await?;
342
343 let ietf_encoding = ietf::Version::try_from(encoding).map_err(|_| Error::Version)?;
345
346 let mut parameters = ietf::Parameters::default();
347 parameters.set_varint(ietf::ParameterVarInt::MaxRequestId, u32::MAX as u64);
348 parameters.set_bytes(ietf::ParameterBytes::Implementation, b"moq-lite-rs".to_vec());
349 if let Some(path) = &self.setup_path {
351 parameters.set_bytes(ietf::ParameterBytes::Path, path.clone().into_bytes());
352 }
353 if let Some(authority) = &self.setup_authority {
354 parameters.set_bytes(ietf::ParameterBytes::Authority, authority.clone().into_bytes());
355 }
356 ietf::solicit::into_setup(&mut parameters, ietf_encoding);
357 ietf::hidden::into_setup(&mut parameters, ietf_encoding);
358 let parameters = parameters.encode_bytes(ietf_encoding)?;
359
360 let client = setup::Client {
361 versions: supported.clone().into(),
362 parameters,
363 };
364
365 stream.writer.encode(&client).await?;
366
367 let mut server: setup::Server = stream.reader.decode().await?;
368
369 let version = supported
370 .iter()
371 .find(|v| coding::Version::from(**v) == server.version)
372 .copied()
373 .ok_or(Error::Version)?;
374
375 let (recv_bw, protocol, goaway) = match version {
376 Version::Lite(v) => {
377 let stream = stream.with_version(v);
378 let start = lite::start(lite::Config {
379 runtime: runtime.clone(),
380 session: session.clone(),
381 setup_stream: Some(stream),
382 publish: publish.clone(),
383 subscribe: subscribe.clone(),
384 peer_hop: self.peer_hop,
385 version: v,
386 our_setup: lite::Setup::default(),
389 peer_setup: None,
390 })?;
391
392 (
393 start.recv_bandwidth,
394 crate::driver::Protocol::Lite(Box::new(start.driver)),
395 start.goaway,
396 )
397 }
398 Version::Ietf(v) => {
399 let parameters = ietf::Parameters::decode(&mut server.parameters, v)?;
402 let request_id_max = parameters
403 .get_varint(ietf::ParameterVarInt::MaxRequestId)
404 .map(ietf::RequestId);
405 let peer_declared = ietf::peer::Peer {
406 solicit: ietf::solicit::from_setup(¶meters, v)?,
407 hidden: ietf::hidden::from_setup(¶meters, v),
408 ..Default::default()
409 };
410
411 let stream = stream.with_version(v);
412 let (protocol, goaway) = ietf::start(ietf::Config {
414 runtime: runtime.clone(),
415 session: session.clone(),
416 setup: Some(stream),
417 request_id_max,
418 client: true,
419 publish: publish.clone(),
420 subscribe: subscribe.clone(),
421 peer_hop: self.peer_hop,
422 cost: self.cost,
423 version: v,
424 path: None,
425 authority: None,
426 peer_setup_stream: None,
427 peer_declared: Some(peer_declared),
428 early_unis: Vec::new(),
429 })?;
430 (None, crate::driver::Protocol::Ietf(protocol), goaway)
431 }
432 };
433
434 Ok(Session::new(runtime, session, version, recv_bw, protocol, goaway))
435 }
436}
437
438#[cfg(test)]
439mod tests {
440 use super::*;
441 use crate::model::ProduceTest;
442 use std::{
443 collections::VecDeque,
444 sync::{Arc, Mutex},
445 };
446
447 use std::task::{Context, Poll};
448
449 use crate::SessionError;
450 use crate::coding::{Decode, Encode};
451 use bytes::{BufMut, Bytes};
452
453 #[derive(Debug, Clone, Default)]
454 struct FakeError;
455
456 impl std::fmt::Display for FakeError {
457 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
458 write!(f, "fake transport error")
459 }
460 }
461
462 impl std::error::Error for FakeError {}
463
464 impl web_transport_trait::Error for FakeError {
465 fn session_error(&self) -> Option<(u32, String)> {
466 Some((0, "closed".to_string()))
467 }
468 }
469
470 #[derive(Clone, Default)]
471 struct FakeSession {
472 state: Arc<FakeSessionState>,
473 park: kio::Park,
475 }
476
477 #[derive(Default)]
478 struct FakeSessionState {
479 protocol: Option<&'static str>,
480 control_stream: Mutex<Option<(FakeSendStream, FakeRecvStream)>>,
481 close_events: Mutex<Vec<(u32, String)>>,
482 closed: kio::Fan,
483 control_writes: Arc<Mutex<Vec<u8>>>,
484 send_rate: Mutex<Option<u64>>,
485 bytes_sent: Mutex<Option<u64>>,
486 }
487
488 impl FakeSession {
489 fn new(protocol: Option<&'static str>, server_control_bytes: Vec<u8>) -> Self {
490 let writes = Arc::new(Mutex::new(Vec::new()));
491 let send = FakeSendStream { writes: writes.clone() };
492 let recv = FakeRecvStream {
493 data: VecDeque::from(server_control_bytes),
494 };
495 let state = FakeSessionState {
496 protocol,
497 control_stream: Mutex::new(Some((send, recv))),
498 close_events: Mutex::new(Vec::new()),
499 closed: kio::Fan::default(),
500 control_writes: writes,
501 send_rate: Mutex::new(None),
502 bytes_sent: Mutex::new(None),
503 };
504 Self {
505 state: Arc::new(state),
506 park: kio::Park::default(),
507 }
508 }
509
510 fn set_send_rate(&self, rate: Option<u64>) {
511 *self.state.send_rate.lock().unwrap() = rate;
512 }
513
514 fn set_bytes_sent(&self, bytes: Option<u64>) {
515 *self.state.bytes_sent.lock().unwrap() = bytes;
516 }
517
518 fn control_writes(&self) -> Vec<u8> {
519 self.state.control_writes.lock().unwrap().clone()
520 }
521
522 async fn wait_for_first_close(&self) -> (u32, String) {
523 kio::wait(|waiter| {
524 self.state.closed.register(waiter);
525 match self.state.close_events.lock().unwrap().first().cloned() {
526 Some(close) => std::task::Poll::Ready(close),
527 None => std::task::Poll::Pending,
528 }
529 })
530 .await
531 }
532 }
533
534 impl web_transport_trait::poll::Session for FakeSession {
535 type SendStream = FakeSendStream;
536 type RecvStream = FakeRecvStream;
537 type Error = FakeError;
538
539 fn poll_accept_uni(&mut self, _cx: &mut Context<'_>) -> Poll<Result<Self::RecvStream, Self::Error>> {
540 Poll::Pending
541 }
542
543 fn poll_accept_bi(
544 &mut self,
545 _cx: &mut Context<'_>,
546 ) -> Poll<Result<(Self::SendStream, Self::RecvStream), Self::Error>> {
547 Poll::Pending
548 }
549
550 fn poll_open_bi(
551 &mut self,
552 _cx: &mut Context<'_>,
553 ) -> Poll<Result<(Self::SendStream, Self::RecvStream), Self::Error>> {
554 Poll::Ready(self.state.control_stream.lock().unwrap().take().ok_or(FakeError))
555 }
556
557 fn poll_open_uni(&mut self, _cx: &mut Context<'_>) -> Poll<Result<Self::SendStream, Self::Error>> {
558 Poll::Pending
559 }
560
561 fn poll_send_datagram(&mut self, _cx: &mut Context<'_>, _payload: &[u8]) -> Poll<Result<(), Self::Error>> {
562 Poll::Ready(Ok(()))
563 }
564
565 fn poll_recv_datagram(&mut self, _cx: &mut Context<'_>) -> Poll<Result<Bytes, Self::Error>> {
566 Poll::Pending
567 }
568
569 fn max_datagram_size(&self) -> usize {
570 1200
571 }
572
573 fn protocol(&self) -> Option<&str> {
574 self.state.protocol
575 }
576
577 fn close(&mut self, code: u32, reason: &str) {
578 self.state.close_events.lock().unwrap().push((code, reason.to_string()));
579 self.state.closed.wake();
580 }
581
582 fn poll_closed(&mut self, cx: &mut Context<'_>) -> Poll<Self::Error> {
583 self.state.closed.register(self.park.hold(cx));
585 match self.state.close_events.lock().unwrap().is_empty() {
586 false => Poll::Ready(FakeError),
587 true => Poll::Pending,
588 }
589 }
590
591 fn stats(&self) -> impl web_transport_trait::Stats {
592 FakeStats {
593 send_rate: *self.state.send_rate.lock().unwrap(),
594 bytes_sent: *self.state.bytes_sent.lock().unwrap(),
595 }
596 }
597 }
598
599 struct FakeStats {
600 send_rate: Option<u64>,
601 bytes_sent: Option<u64>,
602 }
603
604 impl web_transport_trait::Stats for FakeStats {
605 fn estimated_send_rate(&self) -> Option<u64> {
606 self.send_rate
607 }
608
609 fn bytes_sent(&self) -> Option<u64> {
610 self.bytes_sent
611 }
612 }
613
614 #[derive(Clone, Default)]
615 struct FakeSendStream {
616 writes: Arc<Mutex<Vec<u8>>>,
617 }
618
619 impl web_transport_trait::poll::SendStream for FakeSendStream {
620 type Error = FakeError;
621
622 fn poll_write(&mut self, _cx: &mut Context<'_>, buf: &[u8]) -> Poll<Result<usize, Self::Error>> {
623 self.writes.lock().unwrap().put_slice(buf);
624 Poll::Ready(Ok(buf.len()))
625 }
626
627 fn set_priority(&mut self, _order: u8) {}
628
629 fn finish(&mut self) -> Result<(), Self::Error> {
630 Ok(())
631 }
632
633 fn reset(&mut self, _code: u32) {}
634
635 fn poll_closed(&mut self, _cx: &mut Context<'_>) -> Poll<Result<(), Self::Error>> {
636 Poll::Ready(Ok(()))
637 }
638 }
639
640 struct FakeRecvStream {
641 data: VecDeque<u8>,
642 }
643
644 impl web_transport_trait::poll::RecvStream for FakeRecvStream {
645 type Error = FakeError;
646
647 fn poll_read(&mut self, _cx: &mut Context<'_>, dst: &mut [u8]) -> Poll<Result<Option<usize>, Self::Error>> {
648 if self.data.is_empty() {
649 return Poll::Ready(Ok(None));
650 }
651
652 let size = dst.len().min(self.data.len());
653 for slot in dst.iter_mut().take(size) {
654 *slot = self.data.pop_front().unwrap();
655 }
656 Poll::Ready(Ok(Some(size)))
657 }
658
659 fn stop(&mut self, _code: u32) {}
660
661 fn poll_closed(&mut self, _cx: &mut Context<'_>) -> Poll<Result<(), Self::Error>> {
662 Poll::Ready(Ok(()))
663 }
664 }
665
666 fn mock_server_setup(negotiated: Version) -> Vec<u8> {
667 let mut encoded = Vec::new();
668 let server = setup::Server {
669 version: negotiated.into(),
670 parameters: Bytes::new(),
671 };
672 server
673 .encode(&mut encoded, Version::Ietf(ietf::Version::Draft14))
674 .unwrap();
675
676 let info = lite::SessionInfo { bitrate: Some(1) };
678 let lite_v = lite::Version::try_from(negotiated).unwrap();
679 info.encode(&mut encoded, lite_v).unwrap();
680
681 encoded
682 }
683
684 async fn run_alpn_lite_fallback_case(protocol: Option<&'static str>) {
685 let fake = FakeSession::new(protocol, mock_server_setup(Version::Lite(lite::Version::Lite01)));
686 let client = Client::new().with_versions(
687 [
688 Version::Lite(lite::Version::Lite03),
689 Version::Lite(lite::Version::Lite02),
690 Version::Lite(lite::Version::Lite01),
691 Version::Ietf(ietf::Version::Draft14),
692 ]
693 .into(),
694 );
695
696 let (_session, driver) = client
698 .connect(tokio::time::Instant::now().into_std(), fake.clone())
699 .await
700 .unwrap();
701 tokio::spawn(crate::time::run(driver));
702
703 let mut setup_bytes = Bytes::from(fake.control_writes());
705 let setup = setup::Client::decode(&mut setup_bytes, Version::Ietf(ietf::Version::Draft14)).unwrap();
706 let advertised: Vec<Version> = setup.versions.iter().map(|v| Version::try_from(*v).unwrap()).collect();
707 assert_eq!(
708 advertised,
709 vec![
710 Version::Lite(lite::Version::Lite02),
711 Version::Lite(lite::Version::Lite01),
712 Version::Ietf(ietf::Version::Draft14),
713 ]
714 );
715
716 let (code, _) = fake.wait_for_first_close().await;
723 assert_ne!(code, SessionError::Version.to_code(), "SessionInfo failed to decode");
725 }
726
727 #[tokio::test(start_paused = true)]
733 async fn connect_does_not_wait_for_the_peer_to_announce() {
734 let gate = kio::Producer::new(true);
736 let transport = crate::lite::test_transport::SinkSession::gated_bi(gate.consume())
737 .with_protocol(crate::version::ALPN_LITE_05);
738
739 let origin = crate::origin::Config::new(crate::Hop::new(1).unwrap()).produce();
741 let client = Client::new()
742 .with_versions([Version::Lite(lite::Version::Lite05)].into())
743 .with_subscriber(origin);
744
745 let (_session, _driver) = tokio::time::timeout(
748 std::time::Duration::from_secs(30),
749 client.connect(tokio::time::Instant::now().into_std(), transport),
750 )
751 .await
752 .expect("connect waited on a peer that never announced")
753 .expect("connect failed");
754 }
755
756 #[tokio::test(start_paused = true)]
759 async fn draft14_setup_carries_the_authority() {
760 let fake = FakeSession::new(Some(ALPN_LITE), mock_server_setup(Version::Lite(lite::Version::Lite01)));
761 let client = Client::new()
762 .with_versions(
763 [
764 Version::Lite(lite::Version::Lite01),
765 Version::Ietf(ietf::Version::Draft14),
766 ]
767 .into(),
768 )
769 .with_path("/anon")
770 .with_authority("relay.example.com:4443");
771
772 let (_session, driver) = client
773 .connect(tokio::time::Instant::now().into_std(), fake.clone())
774 .await
775 .unwrap();
776 tokio::spawn(crate::time::run(driver));
777
778 let mut setup_bytes = Bytes::from(fake.control_writes());
779 let setup = setup::Client::decode(&mut setup_bytes, Version::Ietf(ietf::Version::Draft14)).unwrap();
780 let mut parameters = setup.parameters;
781 let parameters = ietf::Parameters::decode(&mut parameters, ietf::Version::Draft14).unwrap();
782 assert_eq!(
783 parameters.get_bytes(ietf::ParameterBytes::Authority),
784 Some(b"relay.example.com:4443".as_ref())
785 );
786 assert_eq!(
787 parameters.get_bytes(ietf::ParameterBytes::Path),
788 Some(b"/anon".as_ref())
789 );
790 }
791
792 #[tokio::test(start_paused = true)]
793 async fn alpn_lite_falls_back_to_draft14_and_switches_version_post_setup() {
794 run_alpn_lite_fallback_case(Some(ALPN_LITE)).await;
795 }
796
797 #[tokio::test(start_paused = true)]
798 async fn no_alpn_falls_back_to_draft14_and_switches_version_post_setup() {
799 run_alpn_lite_fallback_case(None).await;
800 }
801
802 #[test]
805 fn driver_is_caller_polled_and_holds_no_session() {
806 let fake = FakeSession::new(Some(ALPN_LITE_04), Vec::new());
807 let client = Client::new().with_versions(Version::Lite(lite::Version::Lite04).into());
808
809 let runtime = crate::runtime::Test::new();
810 let (session, mut driver) = futures::executor::block_on(client.connect(runtime.now(), fake.clone())).unwrap();
811 assert_eq!(session.version(), Version::Lite(lite::Version::Lite04));
812
813 assert!(driver.poll(runtime.now(), &kio::Waiter::noop()).is_ok());
815
816 drop(session);
819 assert!(fake.state.close_events.lock().unwrap().is_empty());
820 let _ = driver.poll(runtime.now(), &kio::Waiter::noop());
821 assert_eq!(
822 fake.state.close_events.lock().unwrap()[0].0,
823 SessionError::Cancel.to_code()
824 );
825 }
826
827 #[test]
831 fn session_clones_share_the_close() {
832 let fake = FakeSession::new(Some(ALPN_LITE_04), Vec::new());
833 let client = Client::new().with_versions(Version::Lite(lite::Version::Lite04).into());
834
835 let runtime = crate::runtime::Test::new();
836 let (session, mut driver) = futures::executor::block_on(client.connect(runtime.now(), fake.clone())).unwrap();
837 let clone = session.clone();
838
839 drop(session);
841 assert!(fake.state.close_events.lock().unwrap().is_empty());
842 let _ = driver.poll(runtime.now(), &kio::Waiter::noop());
843 assert!(fake.state.close_events.lock().unwrap().is_empty());
844
845 clone.abort(Error::Cancel);
846 let _ = driver.poll(runtime.now(), &kio::Waiter::noop());
847 assert_eq!(
848 fake.state.close_events.lock().unwrap()[0].0,
849 SessionError::Cancel.to_code()
850 );
851
852 let _ = driver.poll(runtime.now(), &kio::Waiter::noop());
855 futures::executor::block_on(clone.closed());
856
857 let closes = fake.state.close_events.lock().unwrap().len();
859 drop(clone);
860 let _ = driver.poll(runtime.now(), &kio::Waiter::noop());
861 assert_eq!(fake.state.close_events.lock().unwrap().len(), closes);
862 }
863
864 #[test]
868 fn dropped_driver_resolves_closed() {
869 let fake = FakeSession::new(Some(ALPN_LITE_04), Vec::new());
870 let client = Client::new().with_versions(Version::Lite(lite::Version::Lite04).into());
871
872 let runtime = crate::runtime::Test::new();
873 let (session, driver) = futures::executor::block_on(client.connect(runtime.now(), fake.clone())).unwrap();
874
875 drop(driver);
876 assert!(matches!(futures::executor::block_on(session.closed()), Error::Cancel));
877 }
878
879 #[derive(Clone)]
883 struct LocalSession {
884 inner: FakeSession,
885 _local: std::rc::Rc<()>,
886 }
887
888 struct LocalSend {
889 inner: FakeSendStream,
890 _local: std::rc::Rc<()>,
891 }
892
893 struct LocalRecv {
894 inner: FakeRecvStream,
895 _local: std::rc::Rc<()>,
896 }
897
898 impl web_transport_trait::poll::Session for LocalSession {
899 type SendStream = LocalSend;
900 type RecvStream = LocalRecv;
901 type Error = FakeError;
902
903 fn poll_accept_uni(&mut self, cx: &mut Context<'_>) -> Poll<Result<Self::RecvStream, Self::Error>> {
904 self.inner.poll_accept_uni(cx).map_ok(|stream| LocalRecv {
905 inner: stream,
906 _local: self._local.clone(),
907 })
908 }
909
910 fn poll_accept_bi(
911 &mut self,
912 cx: &mut Context<'_>,
913 ) -> Poll<Result<(Self::SendStream, Self::RecvStream), Self::Error>> {
914 self.inner.poll_accept_bi(cx).map_ok(|(send, recv)| {
915 (
916 LocalSend {
917 inner: send,
918 _local: self._local.clone(),
919 },
920 LocalRecv {
921 inner: recv,
922 _local: self._local.clone(),
923 },
924 )
925 })
926 }
927
928 fn poll_open_bi(
929 &mut self,
930 cx: &mut Context<'_>,
931 ) -> Poll<Result<(Self::SendStream, Self::RecvStream), Self::Error>> {
932 self.inner.poll_open_bi(cx).map_ok(|(send, recv)| {
933 (
934 LocalSend {
935 inner: send,
936 _local: self._local.clone(),
937 },
938 LocalRecv {
939 inner: recv,
940 _local: self._local.clone(),
941 },
942 )
943 })
944 }
945
946 fn poll_open_uni(&mut self, cx: &mut Context<'_>) -> Poll<Result<Self::SendStream, Self::Error>> {
947 self.inner.poll_open_uni(cx).map_ok(|stream| LocalSend {
948 inner: stream,
949 _local: self._local.clone(),
950 })
951 }
952
953 fn poll_send_datagram(&mut self, cx: &mut Context<'_>, payload: &[u8]) -> Poll<Result<(), Self::Error>> {
954 self.inner.poll_send_datagram(cx, payload)
955 }
956
957 fn poll_recv_datagram(&mut self, cx: &mut Context<'_>) -> Poll<Result<Bytes, Self::Error>> {
958 self.inner.poll_recv_datagram(cx)
959 }
960
961 fn max_datagram_size(&self) -> usize {
962 self.inner.max_datagram_size()
963 }
964
965 fn protocol(&self) -> Option<&str> {
966 self.inner.protocol()
967 }
968
969 fn close(&mut self, code: u32, reason: &str) {
970 self.inner.close(code, reason);
971 }
972
973 fn poll_closed(&mut self, cx: &mut Context<'_>) -> Poll<Self::Error> {
974 self.inner.poll_closed(cx)
975 }
976
977 fn stats(&self) -> impl web_transport_trait::Stats {
978 self.inner.stats()
979 }
980 }
981
982 impl web_transport_trait::poll::SendStream for LocalSend {
983 type Error = FakeError;
984
985 fn poll_write(&mut self, cx: &mut Context<'_>, buf: &[u8]) -> Poll<Result<usize, Self::Error>> {
986 self.inner.poll_write(cx, buf)
987 }
988
989 fn set_priority(&mut self, order: u8) {
990 self.inner.set_priority(order);
991 }
992
993 fn finish(&mut self) -> Result<(), Self::Error> {
994 self.inner.finish()
995 }
996
997 fn reset(&mut self, code: u32) {
998 self.inner.reset(code);
999 }
1000
1001 fn poll_closed(&mut self, cx: &mut Context<'_>) -> Poll<Result<(), Self::Error>> {
1002 web_transport_trait::poll::SendStream::poll_closed(&mut self.inner, cx)
1003 }
1004 }
1005
1006 impl web_transport_trait::poll::RecvStream for LocalRecv {
1007 type Error = FakeError;
1008
1009 fn poll_read(&mut self, cx: &mut Context<'_>, dst: &mut [u8]) -> Poll<Result<Option<usize>, Self::Error>> {
1010 self.inner.poll_read(cx, dst)
1011 }
1012
1013 fn stop(&mut self, code: u32) {
1014 self.inner.stop(code);
1015 }
1016
1017 fn poll_closed(&mut self, cx: &mut Context<'_>) -> Poll<Result<(), Self::Error>> {
1018 web_transport_trait::poll::RecvStream::poll_closed(&mut self.inner, cx)
1019 }
1020 }
1021
1022 #[test]
1027 fn connect_lite_over_a_send_less_transport() {
1028 let fake = FakeSession::new(Some(ALPN_LITE_04), Vec::new());
1029 let local = LocalSession {
1030 inner: fake.clone(),
1031 _local: std::rc::Rc::new(()),
1032 };
1033 let client = Client::new().with_versions(Version::Lite(lite::Version::Lite04).into());
1034
1035 let runtime = crate::runtime::Test::new();
1036 let (session, mut driver) = futures::executor::block_on(client.connect_lite(runtime.now(), local)).unwrap();
1037 assert!(driver.poll(runtime.now(), &kio::Waiter::noop()).is_ok());
1038
1039 fn assert_send_sync<T: Send + Sync>(_: &T) {}
1040 assert_send_sync(&session);
1041
1042 session.abort(Error::Cancel);
1043 let _ = driver.poll(runtime.now(), &kio::Waiter::noop());
1044 assert_eq!(
1045 fake.state.close_events.lock().unwrap()[0].0,
1046 SessionError::Cancel.to_code()
1047 );
1048 }
1049
1050 #[test]
1053 fn accept_lite_over_a_send_less_transport() {
1054 let fake = FakeSession::new(Some(ALPN_LITE_04), Vec::new());
1055 let local = LocalSession {
1056 inner: fake.clone(),
1057 _local: std::rc::Rc::new(()),
1058 };
1059 let server = crate::Server::new().with_versions(Version::Lite(lite::Version::Lite04).into());
1060
1061 let runtime = crate::runtime::Test::new();
1062 let (session, mut driver) = futures::executor::block_on(server.accept_lite(runtime.now(), local)).unwrap();
1063 assert_eq!(session.version(), Version::Lite(lite::Version::Lite04));
1064 assert!(driver.poll(runtime.now(), &kio::Waiter::noop()).is_ok());
1065
1066 drop(session);
1067 assert!(fake.state.close_events.lock().unwrap().is_empty());
1068 let _ = driver.poll(runtime.now(), &kio::Waiter::noop());
1069 assert_eq!(
1070 fake.state.close_events.lock().unwrap()[0].0,
1071 SessionError::Cancel.to_code()
1072 );
1073 }
1074
1075 #[test]
1078 fn connect_lite_refuses_ietf_alpns() {
1079 let fake = FakeSession::new(Some(ALPN_19), Vec::new());
1080 let local = LocalSession {
1081 inner: fake,
1082 _local: std::rc::Rc::new(()),
1083 };
1084 let client = Client::new();
1085 let runtime = crate::runtime::Test::new();
1086 let result = futures::executor::block_on(client.connect_lite(runtime.now(), local));
1087 assert!(matches!(result, Err(Error::Version)));
1088 }
1089
1090 #[tokio::test(start_paused = true)]
1094 async fn stats_reads_prime_the_sampler() {
1095 let fake = FakeSession::new(Some(ALPN_LITE_04), Vec::new());
1096 fake.set_send_rate(Some(1_000_000));
1097
1098 let client = Client::new().with_versions(Version::Lite(lite::Version::Lite04).into());
1099 let (session, driver) = client
1100 .connect(tokio::time::Instant::now().into_std(), fake.clone())
1101 .await
1102 .unwrap();
1103 tokio::spawn(crate::time::run(driver));
1104
1105 assert_eq!(
1107 session.stats().estimated_send_rate,
1108 Some(crate::bandwidth::Rate::from_bps(1_000_000))
1109 );
1110
1111 fake.set_send_rate(Some(2_000_000));
1114 while session.stats().estimated_send_rate != Some(crate::bandwidth::Rate::from_bps(2_000_000)) {
1115 tokio::time::sleep(std::time::Duration::from_millis(10)).await;
1116 }
1117 }
1118
1119 #[tokio::test(start_paused = true)]
1126 async fn stats_capture_the_final_counters() {
1127 let fake = FakeSession::new(Some(ALPN_LITE_04), Vec::new());
1128 fake.set_send_rate(None);
1129 fake.set_bytes_sent(Some(0));
1130
1131 let client = Client::new().with_versions(Version::Lite(lite::Version::Lite04).into());
1132 let (session, driver) = client
1133 .connect(tokio::time::Instant::now().into_std(), fake.clone())
1134 .await
1135 .unwrap();
1136 tokio::spawn(crate::time::run(driver));
1137 assert!(
1138 session.send_bandwidth().is_none(),
1139 "no send-rate estimate, so nothing samples on its own"
1140 );
1141
1142 fake.set_bytes_sent(Some(4242));
1143
1144 session.abort(Error::Cancel);
1145 session.closed().await;
1146
1147 assert_eq!(
1148 session.stats().bytes_sent,
1149 Some(4242),
1150 "the closing snapshot must carry the session's final counters"
1151 );
1152 }
1153
1154 #[tokio::test(start_paused = true)]
1158 async fn send_bandwidth_samples_while_the_driver_runs() {
1159 let fake = FakeSession::new(Some(ALPN_LITE_04), Vec::new());
1160 fake.set_send_rate(Some(1_000_000));
1161
1162 let client = Client::new().with_versions(Version::Lite(lite::Version::Lite04).into());
1163 let (session, driver) = client
1164 .connect(tokio::time::Instant::now().into_std(), fake.clone())
1165 .await
1166 .unwrap();
1167 tokio::spawn(crate::time::run(driver));
1168
1169 let mut bandwidth = session.send_bandwidth().expect("backend reports an estimate");
1170 assert_eq!(
1171 bandwidth.changed().await.unwrap(),
1172 Some(crate::bandwidth::Rate::from_bps(1_000_000))
1173 );
1174
1175 fake.set_send_rate(Some(2_000_000));
1177 assert_eq!(
1178 bandwidth.changed().await.unwrap(),
1179 Some(crate::bandwidth::Rate::from_bps(2_000_000))
1180 );
1181 }
1182}