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