1use std::{
4 io::{self, IoSlice},
5 marker::PhantomData,
6 pin::Pin,
7 sync::{Arc, Mutex, MutexGuard},
8 task::{Context, Poll, Waker},
9};
10
11use bytes::{Buf, BufMut, BytesMut, buf::UninitSlice};
12use tokio::io::{AsyncRead, AsyncWrite, ReadBuf};
13
14use crate::{
15 Limits,
16 fragment::Kind,
17 fragment::{Flags, FragmentHeader},
18 transport::{AnyRecv, AnySend, RecvFrame, SendFrame},
19 window::{ControlSink, SessionWindow},
20};
21
22fn lock<T>(mutex: &Mutex<T>) -> MutexGuard<'_, T> {
23 mutex
24 .lock()
25 .unwrap_or_else(|poisoned| poisoned.into_inner())
26}
27
28fn register_waker(slot: &mut Option<Waker>, waker: &Waker) {
29 if !slot
30 .as_ref()
31 .is_some_and(|current| current.will_wake(waker))
32 {
33 *slot = Some(waker.clone());
34 }
35}
36
37fn wake(waker: Option<Waker>) {
38 if let Some(waker) = waker {
39 waker.wake();
40 }
41}
42
43#[derive(Clone, Copy, Debug, Eq, PartialEq)]
44pub(crate) enum SendAction {
45 Fragment,
46 Finish,
47 Abort,
48}
49
50#[derive(Clone, Copy, Debug, Eq, PartialEq)]
51enum SendState {
52 Idle,
54 Demand,
57 Granted,
61 Staging,
66 Fragment,
69 FragmentDemand,
72 Finish,
73 FragmentFinish,
74 Abort,
82 Failed,
87}
88
89pub(crate) struct SendShared {
90 token: Option<AnySend<'static>>,
91 kind: Kind,
92 id: u64,
93 max_fragment_size: usize,
94 copy_threshold: usize,
95 written: usize,
99 reserved: usize,
105 session: Arc<SessionWindow>,
110 buffer: BytesMut,
113 state: SendState,
114 error: Option<(io::ErrorKind, String)>,
119 started: bool,
131 writer_waker: Option<Waker>,
132 driver_waker: Option<Waker>,
133}
134
135impl SendShared {
136 pub(crate) fn new(
137 kind: Kind,
138 id: u64,
139 limits: &Limits,
140 session: Arc<SessionWindow>,
141 ) -> Arc<Mutex<Self>> {
142 Arc::new(Mutex::new(Self {
143 token: None,
144 kind,
145 id,
146 max_fragment_size: limits.max_fragment_size,
147 copy_threshold: limits.trailer_send_copy_threshold,
148 written: 0,
149 reserved: 0,
150 session,
151 buffer: BytesMut::new(),
152 state: SendState::Idle,
153 error: None,
154 started: false,
155 writer_waker: None,
156 driver_waker: None,
157 }))
158 }
159
160 pub(crate) fn start(shared: &Mutex<Self>) {
163 let mut inner = lock(shared);
164 if inner.started {
165 return;
166 }
167 inner.started = true;
168 let writer = inner.writer_waker.take();
169 drop(inner);
170 wake(writer);
171 }
172
173 pub(crate) fn poll_action(shared: &Mutex<Self>, cx: &mut Context<'_>) -> Poll<SendAction> {
174 let mut inner = lock(shared);
175 inner.driver_waker.take();
176 match inner.state {
177 SendState::Demand
178 | SendState::Fragment
179 | SendState::FragmentDemand
180 | SendState::FragmentFinish => Poll::Ready(SendAction::Fragment),
181 SendState::Finish => Poll::Ready(SendAction::Finish),
182 SendState::Abort => Poll::Ready(SendAction::Abort),
183 SendState::Idle | SendState::Granted | SendState::Staging => {
184 register_waker(&mut inner.driver_waker, cx.waker());
185 Poll::Pending
186 }
187 SendState::Failed => {
188 register_waker(&mut inner.driver_waker, cx.waker());
195 Poll::Pending
196 }
197 }
198 }
199
200 pub(crate) unsafe fn grant<'a>(
210 shared: &Arc<Mutex<Self>>,
211 token: AnySend<'a>,
212 max_fragment_size: usize,
213 ) -> SendLease<'a> {
214 let token = unsafe { std::mem::transmute::<AnySend<'a>, AnySend<'static>>(token) };
217 let mut inner = lock(shared);
218 assert!(inner.token.is_none());
219 if inner.buffer.is_empty() {
220 inner.state = SendState::Granted;
221 }
222 inner.token = Some(token);
223 inner.max_fragment_size = max_fragment_size;
224 let writer = inner.writer_waker.take();
225 drop(inner);
226 wake(writer);
227 SendLease {
228 shared: shared.clone(),
229 armed: true,
230 _borrow: PhantomData,
231 }
232 }
233
234 pub(crate) async fn wait_fragment(shared: &Mutex<Self>) -> io::Result<SendAction> {
249 let mut needed_drain = false;
250 let mut yielded = false;
251 loop {
252 let outcome = std::future::poll_fn(|cx| {
253 let mut inner = lock(shared);
254 inner.driver_waker.take();
255 if inner.state == SendState::Failed {
256 let (kind, message) = inner.error.clone().expect("error set for Failed");
257 return Poll::Ready(Err(io::Error::new(kind, message)));
258 }
259 if !inner.buffer.is_empty() {
260 needed_drain = true;
261 let result = poll_flush_buffer(&mut inner, cx);
262 match result {
263 Poll::Pending => return Poll::Pending,
264 Poll::Ready(Err(error)) => {
265 inner.state = SendState::Failed;
266 inner.error = Some((error.kind(), error.to_string()));
267 let writer = inner.writer_waker.take();
268 drop(inner);
269 wake(writer);
270 return Poll::Ready(Err(error));
271 }
272 Poll::Ready(Ok(())) => {}
273 }
274 }
275 match inner.state {
276 SendState::Fragment | SendState::FragmentDemand | SendState::FragmentFinish => {
277 Poll::Ready(Ok(Some(SendAction::Fragment)))
278 }
279 SendState::Finish => Poll::Ready(Ok(Some(SendAction::Finish))),
280 SendState::Abort => Poll::Ready(Ok(Some(SendAction::Abort))),
281 SendState::Granted if !yielded => {
282 Poll::Ready(Ok(None))
285 }
286 SendState::Granted => {
287 inner.state = SendState::Staging;
288 register_waker(&mut inner.driver_waker, cx.waker());
289 Poll::Pending
290 }
291 SendState::Idle | SendState::Demand | SendState::Staging => {
292 register_waker(&mut inner.driver_waker, cx.waker());
297 Poll::Pending
298 }
299 SendState::Failed => unreachable!("handled above"),
300 }
301 })
302 .await?;
303 match outcome {
304 Some(result) => return Ok(result),
305 None => {
306 yielded = true;
307 tokio::task::yield_now().await;
308 }
309 }
310 }
311 }
312
313 fn poll_flush(shared: &Mutex<Self>, cx: &mut Context<'_>) -> Poll<io::Result<()>> {
314 let mut inner = lock(shared);
315 inner.writer_waker.take();
316 match inner.state {
317 SendState::Abort | SendState::Failed => {
318 let (kind, message) = inner.error.clone().expect("error set for Abort/Failed");
319 Poll::Ready(Err(io::Error::new(kind, message)))
320 }
321 _ if !inner.buffer.is_empty() => {
322 register_waker(&mut inner.writer_waker, cx.waker());
323 Poll::Pending
324 }
325 _ => Poll::Ready(Ok(())),
326 }
327 }
328
329 fn finish(shared: &Mutex<Self>) {
330 let mut inner = lock(shared);
331 inner.state = match inner.state {
332 SendState::Fragment | SendState::FragmentDemand => SendState::FragmentFinish,
333 SendState::FragmentFinish => SendState::FragmentFinish,
334 aborted @ (SendState::Abort | SendState::Failed) => aborted,
339 _ => SendState::Finish,
340 };
341 let driver = inner.driver_waker.take();
342 let writer = inner.writer_waker.take();
343 drop(inner);
344 wake(driver);
345 wake(writer);
346 }
347
348 fn abandon(shared: &Mutex<Self>) {
350 Self::set_aborted(shared, io::ErrorKind::BrokenPipe, "trailer is closed");
351 }
352
353 pub(crate) fn discard(shared: &Mutex<Self>) {
364 Self::set_aborted(
365 shared,
366 io::ErrorKind::BrokenPipe,
367 "trailer discarded by peer",
368 );
369 }
370
371 fn set_aborted(shared: &Mutex<Self>, kind: io::ErrorKind, message: &str) {
372 let mut inner = lock(shared);
373 if !matches!(
374 inner.state,
375 SendState::Finish | SendState::FragmentFinish | SendState::Failed
376 ) {
377 inner.state = SendState::Abort;
378 inner.error = Some((kind, message.into()));
379 }
380 let driver = inner.driver_waker.take();
381 let writer = inner.writer_waker.take();
382 let session = inner.session.clone();
383 let id = inner.id;
384 drop(inner);
385 session.settle(id);
390 wake(driver);
391 wake(writer);
392 }
393
394 fn reserve(&mut self, want: usize) -> usize {
402 let granted = self.session.debit_up_to(self.id, want);
403 self.written += granted;
404 granted
405 }
406
407 fn unreserve(&mut self, len: usize) {
409 if len == 0 {
410 return;
411 }
412 self.written -= len;
413 self.session.refund(self.id, len);
414 }
415}
416
417const RESERVED_HELD: &str = "Granted/Staging is only reached from Demand, which reserves";
424
425fn park_for_credit(shared: &mut SendShared, cx: &mut Context<'_>) {
434 register_waker(&mut shared.writer_waker, cx.waker());
437 shared.session.park(cx.waker());
438}
439
440fn poll_flush_buffer(shared: &mut SendShared, cx: &mut Context<'_>) -> Poll<io::Result<()>> {
441 let Some(token) = shared.token.as_mut() else {
442 return Poll::Ready(Err(io::Error::other("send lease has no frame token")));
443 };
444 loop {
445 if shared.buffer.is_empty() {
446 break Poll::Ready(Ok(()));
447 }
448 match token.poll_write_once(cx, &shared.buffer) {
449 Poll::Ready(Ok(0)) => break Poll::Ready(Err(io::ErrorKind::WriteZero.into())),
450 Poll::Ready(Ok(n)) => shared.buffer.advance(n),
451 Poll::Ready(Err(error)) => break Poll::Ready(Err(error)),
452 Poll::Pending => break Poll::Pending,
453 }
454 }
455}
456
457pub(crate) struct SendLease<'a> {
458 shared: Arc<Mutex<SendShared>>,
459 armed: bool,
460 _borrow: PhantomData<&'a mut ()>,
461}
462
463impl SendLease<'_> {
464 pub(crate) fn complete(mut self) {
465 let mut shared = lock(&self.shared);
466 shared.token.take();
467 shared.buffer.clear();
468 shared.state = match shared.state {
469 SendState::Fragment | SendState::FragmentDemand => SendState::Idle,
470 SendState::FragmentFinish => SendState::Finish,
471 state => state,
472 };
473 let writer = shared.writer_waker.take();
474 self.armed = false;
475 drop(shared);
476 wake(writer);
477 }
478}
479
480impl Drop for SendLease<'_> {
481 fn drop(&mut self) {
482 if !self.armed {
483 return;
484 }
485 let mut shared = lock(&self.shared);
486 shared.token.take();
487 shared.buffer = BytesMut::new();
488 if shared.state != SendState::Failed {
493 shared.state = SendState::Abort;
494 if shared.error.is_none() {
495 shared.error = Some((
496 io::ErrorKind::ConnectionAborted,
497 "send grant was revoked".into(),
498 ));
499 }
500 }
501 let writer = shared.writer_waker.take();
502 shared.driver_waker.take();
503 drop(shared);
504 wake(writer);
505 }
506}
507
508pub struct TrailerSend<T> {
514 shared: Arc<Mutex<SendShared>>,
515 completion: Option<T>,
516}
517
518impl<T> TrailerSend<T> {
519 pub(crate) fn new(shared: Arc<Mutex<SendShared>>, completion: T) -> Self {
520 Self {
521 shared,
522 completion: Some(completion),
523 }
524 }
525
526 pub fn finish(mut self) -> T {
532 SendShared::finish(&self.shared);
533 self.completion.take().unwrap()
534 }
535}
536
537impl<T: Unpin> AsyncWrite for TrailerSend<T> {
538 fn poll_write(
539 self: Pin<&mut Self>,
540 cx: &mut Context<'_>,
541 buf: &[u8],
542 ) -> Poll<io::Result<usize>> {
543 if buf.is_empty() {
544 return Poll::Ready(Ok(0));
545 }
546 let this = self.get_mut();
547 let mut inner = lock(&this.shared);
548 inner.writer_waker.take();
549 match inner.state {
550 SendState::Finish | SendState::FragmentFinish => Poll::Ready(Err(io::Error::new(
551 io::ErrorKind::BrokenPipe,
552 "trailer is closed",
553 ))),
554 SendState::Abort | SendState::Failed => {
555 let (kind, message) = inner.error.clone().expect("error set for Abort/Failed");
556 Poll::Ready(Err(io::Error::new(kind, message)))
557 }
558 SendState::Fragment | SendState::FragmentDemand => {
559 inner.state = SendState::FragmentDemand;
565 register_waker(&mut inner.writer_waker, cx.waker());
566 let driver = inner.driver_waker.take();
567 drop(inner);
568 wake(driver);
569 Poll::Pending
570 }
571 SendState::Idle => {
572 if !inner.started {
576 register_waker(&mut inner.writer_waker, cx.waker());
577 return Poll::Pending;
578 }
579 let want = buf.len().min(inner.max_fragment_size.max(1));
588 let len = inner.reserve(want);
589 if len == 0 {
590 park_for_credit(&mut inner, cx);
591 return Poll::Pending;
592 }
593 if len <= inner.copy_threshold {
594 FragmentHeader {
595 flags: Flags::NONE,
596 kind: inner.kind,
597 id: inner.id,
598 payload_len: len,
599 }
600 .encode_into(&mut inner.buffer);
601 inner.buffer.extend_from_slice(&buf[..len]);
602 inner.state = SendState::Fragment;
603 let driver = inner.driver_waker.take();
604 drop(inner);
605 wake(driver);
606 return Poll::Ready(Ok(len));
607 }
608 inner.reserved = len;
611 inner.state = SendState::Demand;
612 register_waker(&mut inner.writer_waker, cx.waker());
613 let driver = inner.driver_waker.take();
614 drop(inner);
615 wake(driver);
616 Poll::Pending
617 }
618 SendState::Demand => {
619 register_waker(&mut inner.writer_waker, cx.waker());
620 Poll::Pending
621 }
622 SendState::Staging => {
623 debug_assert!(inner.reserved > 0, "{RESERVED_HELD}");
627 let len = buf
628 .len()
629 .min(inner.max_fragment_size.max(1))
630 .min(inner.reserved);
631 let unused = inner.reserved - len;
632 inner.unreserve(unused);
633 inner.reserved = 0;
634 FragmentHeader {
635 flags: Flags::NONE,
636 kind: inner.kind,
637 id: inner.id,
638 payload_len: len,
639 }
640 .encode_into(&mut inner.buffer);
641 inner.buffer.extend_from_slice(&buf[..len]);
642 inner.state = SendState::Fragment;
643 let driver = inner.driver_waker.take();
644 drop(inner);
645 wake(driver);
646 Poll::Ready(Ok(len))
647 }
648 SendState::Granted => {
649 debug_assert!(inner.reserved > 0, "{RESERVED_HELD}");
652 let len = buf
653 .len()
654 .min(inner.max_fragment_size.max(1))
655 .min(inner.reserved);
656 let unused = inner.reserved - len;
661 inner.unreserve(unused);
662 inner.reserved = len;
663 FragmentHeader {
664 flags: Flags::NONE,
665 kind: inner.kind,
666 id: inner.id,
667 payload_len: len,
668 }
669 .encode_into(&mut inner.buffer);
670
671 let header_len = inner.buffer.len();
672 let write_result = {
673 let shared = &mut *inner;
674 let bufs = [IoSlice::new(&shared.buffer), IoSlice::new(&buf[..len])];
675 shared
676 .token
677 .as_mut()
678 .expect("installed send token")
679 .poll_write_vectored_once(cx, &bufs)
680 };
681 match write_result {
682 Poll::Ready(Ok(0)) => {
683 let error = io::Error::from(io::ErrorKind::WriteZero);
684 inner.buffer.clear();
685 inner.unreserve(len);
686 inner.reserved = 0;
687 inner.state = SendState::Failed;
688 inner.error = Some((error.kind(), error.to_string()));
689 let driver = inner.driver_waker.take();
690 drop(inner);
691 wake(driver);
692 Poll::Ready(Err(error))
693 }
694 Poll::Ready(Ok(n)) => {
695 debug_assert!(n <= header_len + len);
696 if n < header_len {
697 inner.buffer.advance(n);
698 inner.buffer.extend_from_slice(&buf[..len]);
699 } else {
700 inner.buffer.clear();
701 inner.buffer.extend_from_slice(&buf[n - header_len..len]);
702 }
703 inner.reserved = 0;
704 inner.state = SendState::Fragment;
705 let driver = inner.driver_waker.take();
706 drop(inner);
707 wake(driver);
708 Poll::Ready(Ok(len))
709 }
710 Poll::Ready(Err(error)) => {
711 inner.buffer.clear();
712 inner.unreserve(len);
713 inner.reserved = 0;
714 inner.state = SendState::Failed;
715 inner.error = Some((error.kind(), error.to_string()));
716 let driver = inner.driver_waker.take();
717 drop(inner);
718 wake(driver);
719 Poll::Ready(Err(error))
720 }
721 Poll::Pending => {
722 inner.buffer.clear();
727 Poll::Pending
728 }
729 }
730 }
731 }
732 }
733
734 fn poll_flush(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<io::Result<()>> {
735 SendShared::poll_flush(&self.get_mut().shared, cx)
736 }
737
738 fn poll_shutdown(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<io::Result<()>> {
739 let this = self.get_mut();
740 SendShared::finish(&this.shared);
741 SendShared::poll_flush(&this.shared, cx)
742 }
743}
744
745impl<T> Drop for TrailerSend<T> {
746 fn drop(&mut self) {
747 if self.completion.is_some() {
748 SendShared::abandon(&self.shared);
749 }
750 }
751}
752
753#[derive(Clone, Copy, Debug, Eq, PartialEq)]
754enum RecvState {
755 Idle,
756 Demand,
759 Reading,
763 Unclaimed,
770 Draining,
776 Fragment,
777 FragmentDemand,
780 Eof,
781 Discard,
782 Failed,
786}
787
788pub(crate) struct RecvShared {
789 token: Option<AnyRecv<'static>>,
790 remaining: usize,
791 stage: BytesMut,
796 copy_threshold: usize,
797 demand_copy_threshold: usize,
798 state: RecvState,
799 error: Option<(io::ErrorKind, String)>,
801 reader_waker: Option<Waker>,
802 driver_waker: Option<Waker>,
803 credit_interval: usize,
807 received: usize,
811 retired: usize,
812 pending: usize,
814 auto_release: bool,
817 sink: Arc<dyn ControlSink>,
820 id: u64,
822 session: Arc<SessionWindow>,
825}
826
827impl RecvShared {
828 pub(crate) fn new(
829 copy_threshold: usize,
830 demand_copy_threshold: usize,
831 credit_interval: usize,
832 session: Arc<SessionWindow>,
833 id: u64,
834 sink: Arc<dyn ControlSink>,
835 ) -> Arc<Mutex<Self>> {
836 Arc::new(Mutex::new(Self {
837 token: None,
838 remaining: 0,
839 stage: BytesMut::new(),
840 copy_threshold,
841 demand_copy_threshold,
842 state: RecvState::Idle,
843 error: None,
844 reader_waker: None,
845 driver_waker: None,
846 credit_interval,
847 received: 0,
848 retired: 0,
849 pending: 0,
850 auto_release: true,
851 sink,
852 id,
853 session,
854 }))
855 }
856
857 pub(crate) fn accept_bytes(shared: &Mutex<Self>, len: usize) -> Option<&'static str> {
866 let mut inner = lock(shared);
867 if inner.state == RecvState::Discard {
871 return None;
872 }
873 if !inner.session.accept_bytes(inner.id, len) {
874 return Some("exceeded the session trailer credit window");
875 }
876 inner.received += len;
877 None
878 }
879
880 fn retire(inner: &mut Self, len: usize) -> Option<(Arc<dyn ControlSink>, u64, u32)> {
891 if inner.state == RecvState::Discard {
896 return None;
897 }
898 inner.pending += len;
899 inner.retired += len;
900 if inner.pending * 2 < inner.credit_interval && !Self::must_flush(inner) {
901 return None;
902 }
903 Self::flush(inner)
904 }
905
906 fn must_flush(inner: &Self) -> bool {
920 inner.state == RecvState::Eof || inner.session.is_exhausted() || Self::is_stalled(inner)
921 }
922
923 fn is_stalled(inner: &Self) -> bool {
926 matches!(inner.state, RecvState::Demand | RecvState::FragmentDemand)
927 && inner.stage.is_empty()
928 }
929
930 fn flush(inner: &mut Self) -> Option<(Arc<dyn ControlSink>, u64, u32)> {
932 if inner.state == RecvState::Discard || inner.pending == 0 {
933 return None;
934 }
935 let count = u32::try_from(inner.pending).unwrap_or(u32::MAX) as usize;
936 inner.pending -= count;
937 inner.session.refund(inner.id, count);
938 Some((inner.sink.clone(), inner.id, count as u32))
939 }
940
941 fn retire_and_emit(shared: &Mutex<Self>, len: usize) {
944 let mut inner = lock(shared);
945 if !inner.auto_release {
946 return;
947 }
948 let emit = Self::retire(&mut inner, len);
949 drop(inner);
950 if let Some((sink, id, count)) = emit {
951 sink.credit(id, count);
952 }
953 }
954
955 pub(crate) unsafe fn grant<'a>(
958 shared: &Arc<Mutex<Self>>,
959 token: AnyRecv<'a>,
960 remaining: usize,
961 ) -> RecvLease<'a> {
962 let token = unsafe { std::mem::transmute::<AnyRecv<'a>, AnyRecv<'static>>(token) };
965 let mut inner = lock(shared);
966 assert!(inner.token.is_none());
967 if inner.state != RecvState::Discard {
968 let demanded = inner.state == RecvState::Demand;
969 let copy_threshold = if demanded {
970 inner.demand_copy_threshold
971 } else {
972 inner.copy_threshold
973 };
974 inner.state = if remaining == 0 {
975 if demanded {
976 RecvState::FragmentDemand
977 } else {
978 RecvState::Fragment
979 }
980 } else if remaining <= copy_threshold {
981 RecvState::Draining
982 } else {
983 RecvState::Unclaimed
984 };
985 }
986 inner.token = Some(token);
987 inner.remaining = remaining;
988 let reader = if inner.state == RecvState::Unclaimed {
989 inner.reader_waker.take()
990 } else {
991 None
992 };
993 drop(inner);
994 wake(reader);
995 RecvLease {
996 shared: shared.clone(),
997 armed: true,
998 _borrow: PhantomData,
999 }
1000 }
1001
1002 pub(crate) async fn wait_fragment(shared: &Mutex<Self>) -> io::Result<bool> {
1013 let mut grace_given = false;
1018 std::future::poll_fn(|cx| {
1019 let mut inner = lock(shared);
1020 inner.driver_waker.take();
1021 loop {
1022 match inner.state {
1023 RecvState::Fragment | RecvState::FragmentDemand => {
1024 return Poll::Ready(Ok(false));
1025 }
1026 RecvState::Discard if inner.remaining == 0 => return Poll::Ready(Ok(true)),
1027 RecvState::Draining | RecvState::Discard => {}
1028 RecvState::Reading => {
1029 register_waker(&mut inner.driver_waker, cx.waker());
1032 return Poll::Pending;
1033 }
1034 RecvState::Unclaimed => {
1035 if !grace_given {
1036 grace_given = true;
1044 cx.waker().wake_by_ref();
1045 return Poll::Pending;
1046 }
1047 inner.state = RecvState::Draining;
1048 }
1049 RecvState::Idle | RecvState::Demand | RecvState::Eof => {
1050 register_waker(&mut inner.driver_waker, cx.waker());
1051 return Poll::Pending;
1052 }
1053 RecvState::Failed => {
1054 register_waker(&mut inner.driver_waker, cx.waker());
1060 return Poll::Pending;
1061 }
1062 }
1063 let discard = inner.state == RecvState::Discard;
1064 let result = if discard {
1065 let mut sink = [0u8; 8192];
1066 let n = inner.remaining.min(sink.len());
1067 let mut dest = &mut sink[..n];
1068 inner
1069 .token
1070 .as_mut()
1071 .expect("installed receive token")
1072 .poll_read_once(cx, &mut dest)
1073 } else {
1074 let remaining = inner.remaining;
1075 inner.stage.reserve(remaining);
1076 let RecvShared { token, stage, .. } = &mut *inner;
1077 let mut limited = stage.limit(remaining);
1084 token
1085 .as_mut()
1086 .expect("installed receive token")
1087 .poll_read_once(cx, &mut limited)
1088 };
1089 match result {
1090 Poll::Ready(Ok(0)) => {
1091 return Poll::Ready(Err(io::ErrorKind::UnexpectedEof.into()));
1092 }
1093 Poll::Ready(Ok(n)) => {
1094 inner.remaining -= n;
1095 if inner.remaining == 0 && inner.state == RecvState::Draining {
1096 inner.state = RecvState::Fragment;
1097 }
1098 if !discard {
1099 let reader = inner.reader_waker.take();
1100 wake(reader);
1101 }
1102 }
1105 Poll::Ready(Err(error)) => return Poll::Ready(Err(error)),
1106 Poll::Pending => {
1107 register_waker(&mut inner.driver_waker, cx.waker());
1108 return Poll::Pending;
1109 }
1110 }
1111 }
1112 })
1113 .await
1114 }
1115
1116 pub(crate) fn finish(shared: &Mutex<Self>) {
1128 let mut inner = lock(shared);
1129 inner.state = RecvState::Eof;
1130 let emit = Self::retire(&mut inner, 0);
1132 let reader = inner.reader_waker.take();
1133 drop(inner);
1134 wake(reader);
1135 if let Some((sink, id, count)) = emit {
1136 sink.credit(id, count);
1137 }
1138 }
1139
1140 pub(crate) fn fail(shared: &Mutex<Self>, error: io::Error) {
1141 let mut inner = lock(shared);
1142 inner.state = RecvState::Failed;
1143 inner.error = Some((error.kind(), error.to_string()));
1144 let reader = inner.reader_waker.take();
1145 drop(inner);
1146 wake(reader);
1147 }
1148
1149 pub(crate) fn discard(shared: &Mutex<Self>) {
1163 let mut inner = lock(shared);
1164 if inner.state == RecvState::Discard {
1168 return;
1169 }
1170 let ended = inner.state == RecvState::Eof;
1174 if !matches!(inner.state, RecvState::Eof | RecvState::Failed) {
1175 inner.state = RecvState::Discard;
1176 }
1177 let driver = inner.driver_waker.take();
1178 let notify = (!ended).then(|| (inner.sink.clone(), inner.id));
1179 inner.pending = 0;
1183 inner.retired = inner.received;
1184 let session = inner.session.clone();
1185 let id = inner.id;
1186 drop(inner);
1187 session.settle(id);
1188 wake(driver);
1189 if let Some((sink, id)) = notify {
1190 sink.discard(id);
1191 }
1192 }
1193}
1194
1195pub(crate) struct RecvLease<'a> {
1196 shared: Arc<Mutex<RecvShared>>,
1197 armed: bool,
1198 _borrow: PhantomData<&'a mut ()>,
1199}
1200
1201impl RecvLease<'_> {
1202 pub(crate) fn complete(mut self) {
1203 let mut shared = lock(&self.shared);
1204 shared.token.take();
1205 shared.remaining = 0;
1206 shared.state = match shared.state {
1207 RecvState::Fragment => RecvState::Idle,
1208 RecvState::FragmentDemand => RecvState::Demand,
1209 state => state,
1210 };
1211 self.armed = false;
1212 }
1213}
1214
1215impl Drop for RecvLease<'_> {
1216 fn drop(&mut self) {
1217 if !self.armed {
1218 return;
1219 }
1220 let mut shared = lock(&self.shared);
1221 shared.token.take();
1222 shared.remaining = 0;
1223 shared.state = RecvState::Failed;
1227 if shared.error.is_none() {
1228 shared.error = Some((
1229 io::ErrorKind::ConnectionAborted,
1230 "receive grant was revoked".into(),
1231 ));
1232 }
1233 let reader = shared.reader_waker.take();
1234 shared.driver_waker.take();
1235 drop(shared);
1236 wake(reader);
1237 }
1238}
1239
1240pub struct TrailerRecv {
1242 pub(crate) shared: Arc<Mutex<RecvShared>>,
1243}
1244
1245impl TrailerRecv {
1246 pub(crate) fn new(shared: Arc<Mutex<RecvShared>>) -> Self {
1247 Self { shared }
1248 }
1249
1250 pub(crate) fn set_manual_credit(&mut self) {
1258 let mut inner = lock(&self.shared);
1259 debug_assert!(
1260 inner.retired == 0,
1261 "manual credit must be selected before the first read"
1262 );
1263 inner.auto_release = false;
1264 }
1265
1266 pub fn release(&mut self, n: usize) {
1286 let mut inner = lock(&self.shared);
1287 debug_assert!(
1288 !inner.auto_release,
1289 "release requires a manual-credit trailer; credit is returned on read otherwise"
1290 );
1291 debug_assert!(
1292 inner.retired + n <= inner.received,
1293 "released more trailer credit than was delivered"
1294 );
1295 if n == 0 {
1296 return;
1297 }
1298 let emit = RecvShared::retire(&mut inner, n);
1299 drop(inner);
1300 if let Some((sink, id, count)) = emit {
1301 sink.credit(id, count);
1302 }
1303 }
1304}
1305
1306impl AsyncRead for TrailerRecv {
1307 fn poll_read(
1308 self: Pin<&mut Self>,
1309 cx: &mut Context<'_>,
1310 buf: &mut ReadBuf<'_>,
1311 ) -> Poll<io::Result<()>> {
1312 let this = self.get_mut();
1313 if buf.remaining() == 0 {
1314 return Poll::Ready(Ok(()));
1315 }
1316 let mut inner = lock(&this.shared);
1317 inner.reader_waker.take();
1318 if !inner.stage.is_empty() {
1333 let n = buf.remaining().min(inner.stage.len());
1334 buf.put_slice(&inner.stage[..n]);
1335 let _ = inner.stage.split_to(n);
1336 drop(inner);
1341 RecvShared::retire_and_emit(&this.shared, n);
1342 return Poll::Ready(Ok(()));
1343 }
1344 match inner.state {
1345 RecvState::Failed => {
1346 let (kind, message) = inner.error.clone().expect("error set for Failed");
1347 Poll::Ready(Err(io::Error::new(kind, message)))
1348 }
1349 RecvState::Eof => Poll::Ready(Ok(())),
1350 RecvState::Idle | RecvState::Fragment => {
1351 inner.state = if inner.state == RecvState::Idle {
1352 RecvState::Demand
1353 } else {
1354 RecvState::FragmentDemand
1355 };
1356 register_waker(&mut inner.reader_waker, cx.waker());
1357 let emit = RecvShared::flush(&mut inner);
1366 drop(inner);
1367 if let Some((sink, id, count)) = emit {
1368 sink.credit(id, count);
1369 }
1370 Poll::Pending
1371 }
1372 RecvState::Demand | RecvState::FragmentDemand => {
1373 register_waker(&mut inner.reader_waker, cx.waker());
1374 Poll::Pending
1375 }
1376 RecvState::Draining | RecvState::Discard => {
1377 register_waker(&mut inner.reader_waker, cx.waker());
1381 Poll::Pending
1382 }
1383 RecvState::Reading | RecvState::Unclaimed => {
1384 let before = buf.filled().len();
1388 let mut adapter = ReadBufMut(buf);
1389 let mut limited = (&mut adapter).limit(inner.remaining);
1390 let result = inner
1391 .token
1392 .as_mut()
1393 .expect("installed receive token")
1394 .poll_read_once(cx, &mut limited);
1395 match result {
1396 Poll::Ready(Ok(0)) => Poll::Ready(Err(io::ErrorKind::UnexpectedEof.into())),
1397 Poll::Ready(Ok(n)) => {
1398 inner.remaining -= n;
1399 if inner.remaining == 0 {
1400 inner.state = RecvState::Fragment;
1401 } else {
1402 inner.state = RecvState::Draining;
1421 }
1422 let driver = inner.driver_waker.take();
1423 drop(inner);
1424 wake(driver);
1425 RecvShared::retire_and_emit(&this.shared, n);
1427 debug_assert_eq!(buf.filled().len() - before, n);
1428 Poll::Ready(Ok(()))
1429 }
1430 Poll::Ready(Err(error)) => Poll::Ready(Err(error)),
1431 Poll::Pending => {
1432 inner.state = RecvState::Reading;
1437 Poll::Pending
1438 }
1439 }
1440 }
1441 }
1442 }
1443}
1444
1445impl Drop for TrailerRecv {
1446 fn drop(&mut self) {
1447 RecvShared::discard(&self.shared);
1448 }
1449}
1450
1451struct ReadBufMut<'a, 'b>(&'a mut ReadBuf<'b>);
1452
1453unsafe impl BufMut for ReadBufMut<'_, '_> {
1454 fn remaining_mut(&self) -> usize {
1455 self.0.remaining()
1456 }
1457
1458 unsafe fn advance_mut(&mut self, cnt: usize) {
1459 unsafe { self.0.assume_init(cnt) };
1461 self.0.advance(cnt);
1462 }
1463
1464 fn chunk_mut(&mut self) -> &mut UninitSlice {
1465 let unfilled = unsafe { self.0.unfilled_mut() };
1467 unsafe { UninitSlice::from_raw_parts_mut(unfilled.as_mut_ptr().cast(), unfilled.len()) }
1470 }
1471}
1472
1473#[cfg(test)]
1474mod tests {
1475 use super::*;
1476 use crate::transport::{AnyReceiver, AnySender, Receiver, Sender, generic};
1477
1478 fn unbounded_limits() -> Limits {
1481 Limits {
1482 trailer_session_window: usize::MAX,
1483 ..zero_copy_limits()
1484 }
1485 }
1486
1487 fn zero_copy_limits() -> Limits {
1491 Limits {
1492 trailer_send_copy_threshold: 0,
1493 ..Limits::default()
1494 }
1495 }
1496
1497 fn demand<T: Unpin>(trailer: &mut TrailerSend<T>, buf: &[u8]) {
1502 let mut cx = Context::from_waker(Waker::noop());
1503 assert!(
1504 Pin::new(trailer).poll_write(&mut cx, buf).is_pending(),
1505 "a write past the copy threshold must demand a token"
1506 );
1507 }
1508
1509 fn send_shared(limits: Limits) -> Arc<Mutex<SendShared>> {
1510 send_shared_id(1, limits)
1511 }
1512
1513 fn send_shared_id(id: u64, limits: Limits) -> Arc<Mutex<SendShared>> {
1514 let session = Arc::new(SessionWindow::new(limits.trailer_session_window));
1515 let shared = SendShared::new(Kind::Request, id, &limits, session);
1516 SendShared::start(&shared);
1519 shared
1520 }
1521
1522 #[derive(Default)]
1525 struct RecordingSink {
1526 credits: Mutex<Vec<(u64, u32)>>,
1527 discards: Mutex<Vec<u64>>,
1528 }
1529
1530 impl ControlSink for Arc<RecordingSink> {
1531 fn payload_credit(&self, _count: u32) {
1532 unreachable!("trailer tests never release payload quota")
1533 }
1534
1535 fn credit(&self, id: u64, count: u32) {
1536 lock(&self.credits).push((id, count));
1537 }
1538
1539 fn discard(&self, id: u64) {
1540 lock(&self.discards).push(id);
1541 }
1542 }
1543
1544 fn recv_shared(limits: Limits) -> (Arc<Mutex<RecvShared>>, Arc<RecordingSink>) {
1545 let session = Arc::new(SessionWindow::new(limits.trailer_session_window));
1546 let sink = Arc::new(RecordingSink::default());
1547 let shared = RecvShared::new(
1548 limits.trailer_recv_copy_threshold,
1549 limits.trailer_recv_demand_copy_threshold,
1550 limits.trailer_credit_interval,
1551 session,
1552 7,
1553 Arc::new(sink.clone()),
1554 );
1555 (shared, sink)
1556 }
1557
1558 fn poll_read_once(trailer: &mut TrailerRecv, output: &mut [u8]) -> Poll<io::Result<usize>> {
1559 let mut read = ReadBuf::new(output);
1560 let mut cx = Context::from_waker(Waker::noop());
1561 match Pin::new(trailer).poll_read(&mut cx, &mut read) {
1562 Poll::Ready(Ok(())) => Poll::Ready(Ok(read.filled().len())),
1563 Poll::Ready(Err(error)) => Poll::Ready(Err(error)),
1564 Poll::Pending => Poll::Pending,
1565 }
1566 }
1567
1568 struct CappedSink {
1569 bytes: Arc<Mutex<Vec<u8>>>,
1570 max_write: usize,
1571 }
1572
1573 impl AsyncWrite for CappedSink {
1574 fn poll_write(
1575 self: Pin<&mut Self>,
1576 _cx: &mut Context<'_>,
1577 buf: &[u8],
1578 ) -> Poll<io::Result<usize>> {
1579 let len = buf.len().min(self.max_write);
1580 lock(&self.bytes).extend_from_slice(&buf[..len]);
1581 Poll::Ready(Ok(len))
1582 }
1583
1584 fn poll_write_vectored(
1585 self: Pin<&mut Self>,
1586 _cx: &mut Context<'_>,
1587 bufs: &[IoSlice<'_>],
1588 ) -> Poll<io::Result<usize>> {
1589 let mut remaining = self.max_write;
1590 let mut written = 0;
1591 let mut output = lock(&self.bytes);
1592 for buf in bufs {
1593 let len = buf.len().min(remaining);
1594 output.extend_from_slice(&buf[..len]);
1595 written += len;
1596 remaining -= len;
1597 if remaining == 0 {
1598 break;
1599 }
1600 }
1601 Poll::Ready(Ok(written))
1602 }
1603
1604 fn poll_flush(self: Pin<&mut Self>, _cx: &mut Context<'_>) -> Poll<io::Result<()>> {
1605 Poll::Ready(Ok(()))
1606 }
1607
1608 fn poll_shutdown(self: Pin<&mut Self>, _cx: &mut Context<'_>) -> Poll<io::Result<()>> {
1609 Poll::Ready(Ok(()))
1610 }
1611 }
1612
1613 #[test]
1614 fn small_send_stages_without_a_grant_and_large_send_demands_one() {
1615 let limits = Limits {
1616 max_fragment_size: 8,
1617 trailer_send_copy_threshold: 4,
1618 ..Limits::default()
1619 };
1620
1621 let small_shared = send_shared(limits);
1622 let mut small = TrailerSend::new(small_shared.clone(), ());
1623 let mut cx = Context::from_waker(Waker::noop());
1624 assert!(matches!(
1625 Pin::new(&mut small).poll_write(&mut cx, b"data"),
1626 Poll::Ready(Ok(4))
1627 ));
1628 let header = FragmentHeader {
1629 flags: Flags::NONE,
1630 kind: Kind::Request,
1631 id: 1,
1632 payload_len: 4,
1633 }
1634 .encode();
1635 assert_eq!(
1636 &small_shared.lock().unwrap().buffer[..],
1637 [&header[..], b"data"].concat()
1638 );
1639 assert_eq!(small_shared.lock().unwrap().state, SendState::Fragment);
1640
1641 let large_shared = send_shared(limits);
1642 let mut large = TrailerSend::new(large_shared.clone(), ());
1643 assert!(
1644 Pin::new(&mut large)
1645 .poll_write(&mut cx, b"large")
1646 .is_pending()
1647 );
1648 assert_eq!(large_shared.lock().unwrap().state, SendState::Demand);
1649 assert!(large_shared.lock().unwrap().buffer.is_empty());
1650 }
1651
1652 #[test]
1653 fn receive_copy_threshold_depends_on_demand_for_this_fragment() {
1654 let undemanded = recv_shared(Limits {
1655 trailer_recv_copy_threshold: 1,
1656 trailer_recv_demand_copy_threshold: 4,
1657 ..unbounded_limits()
1658 })
1659 .0;
1660 let (_, receiver) = generic(tokio::io::empty(), tokio::io::sink());
1661 let mut receiver = AnyReceiver::Generic(receiver);
1662 let lease = unsafe { RecvShared::grant(&undemanded, receiver.recv(), 4) };
1663 assert_eq!(undemanded.lock().unwrap().state, RecvState::Unclaimed);
1664 drop(lease);
1665
1666 let demanded = recv_shared(Limits {
1667 trailer_recv_copy_threshold: 1,
1668 trailer_recv_demand_copy_threshold: 4,
1669 ..unbounded_limits()
1670 })
1671 .0;
1672 let mut trailer = TrailerRecv::new(demanded.clone());
1673 let mut output = [0; 4];
1674 assert!(poll_read_once(&mut trailer, &mut output).is_pending());
1675 assert_eq!(demanded.lock().unwrap().state, RecvState::Demand);
1676 let (_, receiver) = generic(tokio::io::empty(), tokio::io::sink());
1677 let mut receiver = AnyReceiver::Generic(receiver);
1678 let lease = unsafe { RecvShared::grant(&demanded, receiver.recv(), 4) };
1679 assert_eq!(demanded.lock().unwrap().state, RecvState::Draining);
1680 drop(lease);
1681 }
1682
1683 #[test]
1684 fn demand_at_a_completed_fragment_boundary_applies_to_the_next_fragment() {
1685 let shared = recv_shared(Limits {
1686 trailer_recv_copy_threshold: 0,
1687 trailer_recv_demand_copy_threshold: 0,
1688 ..unbounded_limits()
1689 })
1690 .0;
1691 let mut trailer = TrailerRecv::new(shared.clone());
1692 let (_, receiver) = generic(tokio::io::empty(), tokio::io::sink());
1693 let mut receiver = AnyReceiver::Generic(receiver);
1694 let lease = unsafe { RecvShared::grant(&shared, receiver.recv(), 0) };
1695 assert_eq!(shared.lock().unwrap().state, RecvState::Fragment);
1696
1697 let mut output = [0; 1];
1698 assert!(poll_read_once(&mut trailer, &mut output).is_pending());
1699 assert_eq!(shared.lock().unwrap().state, RecvState::FragmentDemand);
1700 lease.complete();
1701 assert_eq!(shared.lock().unwrap().state, RecvState::Demand);
1702 }
1703
1704 #[tokio::test]
1705 async fn unclaimed_large_receive_falls_back_to_driver_draining() {
1706 use tokio::io::AsyncWriteExt;
1707
1708 let shared = recv_shared(Limits {
1709 trailer_recv_copy_threshold: 0,
1710 trailer_recv_demand_copy_threshold: 0,
1711 ..unbounded_limits()
1712 })
1713 .0;
1714 let (mut writer, reader) = tokio::io::duplex(16);
1715 writer.write_all(b"data").await.unwrap();
1716 let (_, receiver) = generic(reader, tokio::io::sink());
1717 let mut receiver = AnyReceiver::Generic(receiver);
1718 let lease = unsafe { RecvShared::grant(&shared, receiver.recv(), 4) };
1719 assert_eq!(shared.lock().unwrap().state, RecvState::Unclaimed);
1720
1721 assert!(!RecvShared::wait_fragment(&shared).await.unwrap());
1722 assert_eq!(shared.lock().unwrap().state, RecvState::Fragment);
1723 assert_eq!(&shared.lock().unwrap().stage[..], b"data");
1724 lease.complete();
1725 }
1726
1727 #[tokio::test]
1728 async fn demanded_large_receive_can_claim_the_grant_directly() {
1729 use tokio::io::AsyncWriteExt;
1730
1731 let shared = recv_shared(Limits {
1732 trailer_recv_copy_threshold: 0,
1733 trailer_recv_demand_copy_threshold: 0,
1734 ..unbounded_limits()
1735 })
1736 .0;
1737 let mut trailer = TrailerRecv::new(shared.clone());
1738 let mut output = [0; 4];
1739 assert!(poll_read_once(&mut trailer, &mut output).is_pending());
1740
1741 let (mut writer, reader) = tokio::io::duplex(16);
1742 writer.write_all(b"data").await.unwrap();
1743 let (_, receiver) = generic(reader, tokio::io::sink());
1744 let mut receiver = AnyReceiver::Generic(receiver);
1745 let lease = unsafe { RecvShared::grant(&shared, receiver.recv(), 4) };
1746 assert_eq!(shared.lock().unwrap().state, RecvState::Unclaimed);
1747 assert!(matches!(
1748 poll_read_once(&mut trailer, &mut output),
1749 Poll::Ready(Ok(4))
1750 ));
1751 assert_eq!(&output, b"data");
1752 assert!(!RecvShared::wait_fragment(&shared).await.unwrap());
1753 lease.complete();
1754 assert_eq!(shared.lock().unwrap().state, RecvState::Idle);
1755 }
1756
1757 #[tokio::test]
1758 async fn abandoned_fragment_flushes_only_its_real_staged_suffix() {
1759 let output = Arc::new(Mutex::new(Vec::new()));
1760 let (sender, _) = generic(
1761 tokio::io::empty(),
1762 CappedSink {
1763 bytes: output.clone(),
1764 max_write: 16,
1765 },
1766 );
1767 let mut sender = AnySender::Generic(sender);
1768 let shared = send_shared(unbounded_limits());
1769 let data = (0..100).map(|value| value as u8).collect::<Vec<_>>();
1770 let mut trailer = TrailerSend::new(shared.clone(), ());
1771 demand(&mut trailer, &data);
1772 let lease = unsafe { SendShared::grant(&shared, sender.send(), 1024) };
1773
1774 let written = std::future::poll_fn(|cx| Pin::new(&mut trailer).poll_write(cx, &data))
1775 .await
1776 .unwrap();
1777 assert_eq!(written, data.len());
1778 {
1779 let inner = lock(&shared);
1780 let header_len = FragmentHeader {
1781 flags: Flags::NONE,
1782 kind: Kind::Request,
1783 id: 1,
1784 payload_len: data.len(),
1785 }
1786 .encode()
1787 .len();
1788 assert_eq!(&inner.buffer[..], &data[16 - header_len..]);
1789 }
1790
1791 drop(trailer);
1792 assert_eq!(
1793 SendShared::wait_fragment(&shared).await.unwrap(),
1794 SendAction::Abort
1795 );
1796 let header_len = FragmentHeader {
1797 flags: Flags::NONE,
1798 kind: Kind::Request,
1799 id: 1,
1800 payload_len: data.len(),
1801 }
1802 .encode()
1803 .len();
1804 assert_eq!(&lock(&output)[header_len..], data);
1805 lease.complete();
1806 }
1807
1808 #[tokio::test]
1809 async fn partial_header_and_payload_share_the_stage_buffer() {
1810 let output = Arc::new(Mutex::new(Vec::new()));
1811 let (sender, _) = generic(
1812 tokio::io::empty(),
1813 CappedSink {
1814 bytes: output.clone(),
1815 max_write: 5,
1816 },
1817 );
1818 let mut sender = AnySender::Generic(sender);
1819 let shared = send_shared_id(7, unbounded_limits());
1820 let data = (0..32).map(|value| value as u8).collect::<Vec<_>>();
1821 let mut trailer = TrailerSend::new(shared.clone(), ());
1822 demand(&mut trailer, &data);
1823 let lease = unsafe { SendShared::grant(&shared, sender.send(), 1024) };
1824
1825 let written = std::future::poll_fn(|cx| Pin::new(&mut trailer).poll_write(cx, &data))
1826 .await
1827 .unwrap();
1828 assert_eq!(written, data.len());
1829
1830 let header = FragmentHeader {
1831 flags: Flags::NONE,
1832 kind: Kind::Request,
1833 id: 7,
1834 payload_len: data.len(),
1835 }
1836 .encode();
1837 let mut expected_stage = Vec::from(&header[5..]);
1838 expected_stage.extend_from_slice(&data);
1839 assert_eq!(&lock(&shared).buffer[..], expected_stage);
1840
1841 assert_eq!(
1842 SendShared::wait_fragment(&shared).await.unwrap(),
1843 SendAction::Fragment
1844 );
1845 assert_eq!(&lock(&output)[..], [&header[..], &data].concat());
1846 lease.complete();
1847 }
1848
1849 #[tokio::test]
1850 async fn finish_releases_an_unused_live_grant() {
1851 let (sender, _) = generic(tokio::io::empty(), tokio::io::sink());
1852 let mut sender = AnySender::Generic(sender);
1853 let shared = send_shared(unbounded_limits());
1854 let lease = unsafe { SendShared::grant(&shared, sender.send(), 1024) };
1855
1856 TrailerSend::new(shared.clone(), ()).finish();
1857 assert_eq!(
1858 SendShared::wait_fragment(&shared).await.unwrap(),
1859 SendAction::Finish
1860 );
1861 lease.complete();
1862 }
1863
1864 #[tokio::test]
1873 async fn an_unstarted_trailer_reserves_no_credit() {
1874 let session = Arc::new(SessionWindow::new(64));
1875 let shared = SendShared::new(Kind::Request, 1, &Limits::default(), session.clone());
1879 let mut trailer = TrailerSend::new(shared.clone(), ());
1880
1881 let mut cx = Context::from_waker(Waker::noop());
1882 assert!(matches!(
1883 Pin::new(&mut trailer).poll_write(&mut cx, b"abcd"),
1884 Poll::Pending
1885 ));
1886 assert_eq!(session.available(), 64, "nothing may be spent yet");
1887 assert_eq!(lock(&shared).state, SendState::Idle);
1888
1889 SendShared::start(&shared);
1890 let written = std::future::poll_fn(|cx| Pin::new(&mut trailer).poll_write(cx, b"abcd"))
1891 .await
1892 .unwrap();
1893 assert_eq!(written, 4);
1894 assert_eq!(session.available(), 60);
1895 }
1896
1897 #[tokio::test]
1898 async fn exhausted_credit_parks_the_writer_until_credit_arrives() {
1899 let (sender, _) = generic(tokio::io::empty(), tokio::io::sink());
1900 let mut sender = AnySender::Generic(sender);
1901 let session = Arc::new(SessionWindow::new(4));
1902 let shared = SendShared::new(Kind::Request, 1, &zero_copy_limits(), session.clone());
1903 SendShared::start(&shared);
1904 let mut trailer = TrailerSend::new(shared.clone(), ());
1905
1906 demand(&mut trailer, b"abcd");
1908 let lease = unsafe { SendShared::grant(&shared, sender.send(), 1024) };
1909 let written = std::future::poll_fn(|cx| Pin::new(&mut trailer).poll_write(cx, b"abcd"))
1910 .await
1911 .unwrap();
1912 assert_eq!(written, 4);
1913 let action = SendShared::wait_fragment(&shared).await.unwrap();
1914 assert_eq!(action, SendAction::Fragment);
1915 lease.complete();
1916
1917 let mut cx = Context::from_waker(Waker::noop());
1920 assert!(matches!(
1921 Pin::new(&mut trailer).poll_write(&mut cx, b"e"),
1922 Poll::Pending
1923 ));
1924 assert_eq!(lock(&shared).state, SendState::Idle);
1925 assert!(matches!(
1926 SendShared::poll_action(&shared, &mut cx),
1927 Poll::Pending
1928 ));
1929
1930 session.refund(1, 2);
1932 demand(&mut trailer, b"ef");
1933 let lease = unsafe { SendShared::grant(&shared, sender.send(), 1024) };
1934 let written = std::future::poll_fn(|cx| Pin::new(&mut trailer).poll_write(cx, b"ef"))
1935 .await
1936 .unwrap();
1937 assert_eq!(written, 2);
1938 lease.complete();
1939 }
1940
1941 #[tokio::test]
1945 async fn pool_below_fragment_size_still_makes_progress() {
1946 let (sender, _) = generic(tokio::io::empty(), tokio::io::sink());
1947 let mut sender = AnySender::Generic(sender);
1948 let session = Arc::new(SessionWindow::new(3));
1949 let limits = Limits {
1950 max_fragment_size: 1024,
1951 ..zero_copy_limits()
1952 };
1953 let shared = SendShared::new(Kind::Request, 1, &limits, session.clone());
1954 SendShared::start(&shared);
1955 let mut trailer = TrailerSend::new(shared.clone(), ());
1956 let mut total = 0;
1957 for _ in 0..4 {
1958 demand(&mut trailer, b"abcdefghij");
1959 let lease = unsafe { SendShared::grant(&shared, sender.send(), 1024) };
1960 let n = std::future::poll_fn(|cx| Pin::new(&mut trailer).poll_write(cx, b"abcdefghij"))
1961 .await
1962 .unwrap();
1963 assert_eq!(n, 3, "each write is clamped to the pool, not dropped");
1964 total += n;
1965 SendShared::wait_fragment(&shared).await.unwrap();
1966 lease.complete();
1967 session.refund(1, 3);
1968 }
1969 assert_eq!(total, 12);
1970 }
1971
1972 #[tokio::test]
1977 async fn one_trailer_exhausting_the_pool_parks_another() {
1978 let (sender, _) = generic(tokio::io::empty(), tokio::io::sink());
1979 let mut sender = AnySender::Generic(sender);
1980 let session = Arc::new(SessionWindow::new(4));
1981 let limits = zero_copy_limits();
1982 let first = SendShared::new(Kind::Request, 1, &limits, session.clone());
1983 let second = SendShared::new(Kind::Request, 2, &limits, session.clone());
1984 SendShared::start(&first);
1985 SendShared::start(&second);
1986
1987 let mut trailer = TrailerSend::new(first.clone(), ());
1988 demand(&mut trailer, b"abcdefgh");
1989 let lease = unsafe { SendShared::grant(&first, sender.send(), 1024) };
1990 let written = std::future::poll_fn(|cx| Pin::new(&mut trailer).poll_write(cx, b"abcdefgh"))
1991 .await
1992 .unwrap();
1993 assert_eq!(written, 4, "clamped by the pool");
1994 SendShared::wait_fragment(&first).await.unwrap();
1995 lease.complete();
1996
1997 let mut other = TrailerSend::new(second.clone(), ());
1998 let mut cx = Context::from_waker(Waker::noop());
1999 assert!(matches!(
2000 Pin::new(&mut other).poll_write(&mut cx, b"i"),
2001 Poll::Pending
2002 ));
2003 assert_eq!(
2004 lock(&second).state,
2005 SendState::Idle,
2006 "a parked writer must hold no transport grant"
2007 );
2008 assert!(matches!(
2009 SendShared::poll_action(&second, &mut cx),
2010 Poll::Pending
2011 ));
2012
2013 session.refund(1, 2);
2014 demand(&mut other, b"ij");
2015 let lease = unsafe { SendShared::grant(&second, sender.send(), 1024) };
2016 let written = std::future::poll_fn(|cx| Pin::new(&mut other).poll_write(cx, b"ij"))
2017 .await
2018 .unwrap();
2019 assert_eq!(written, 2);
2020 lease.complete();
2021 }
2022
2023 #[tokio::test]
2026 async fn aborting_a_trailer_returns_its_session_debt() {
2027 let (sender, _) = generic(tokio::io::empty(), tokio::io::sink());
2028 let mut sender = AnySender::Generic(sender);
2029 let session = Arc::new(SessionWindow::new(64));
2030 let shared = SendShared::new(Kind::Request, 1, &zero_copy_limits(), session.clone());
2031 SendShared::start(&shared);
2032 let mut trailer = TrailerSend::new(shared.clone(), ());
2033 demand(&mut trailer, b"abcdefgh");
2034 let lease = unsafe { SendShared::grant(&shared, sender.send(), 1024) };
2035 std::future::poll_fn(|cx| Pin::new(&mut trailer).poll_write(cx, b"abcdefgh"))
2036 .await
2037 .unwrap();
2038 SendShared::wait_fragment(&shared).await.unwrap();
2039 lease.complete();
2040 assert_eq!(session.available(), 56);
2041
2042 SendShared::discard(&shared);
2043 assert_eq!(session.available(), 64, "the pool is made whole again");
2044 }
2045
2046 #[test]
2049 fn auto_release_credits_on_delivery() {
2050 let (shared, sink) = recv_shared(Limits {
2051 trailer_credit_interval: 8,
2052 ..Limits::default()
2053 });
2054 RecvShared::accept_bytes(&shared, 8);
2055 lock(&shared).stage.extend_from_slice(b"abcdefgh");
2056 let mut trailer = TrailerRecv::new(shared.clone());
2057
2058 let mut out = [0u8; 8];
2059 assert!(matches!(
2060 poll_read_once(&mut trailer, &mut out),
2061 Poll::Ready(Ok(8))
2062 ));
2063 assert_eq!(&*lock(&sink.credits), &[(7, 8)]);
2064 }
2065
2066 #[test]
2068 fn manual_release_does_not_credit_on_read() {
2069 let (shared, sink) = recv_shared(Limits {
2070 trailer_credit_interval: 8,
2071 ..Limits::default()
2072 });
2073 RecvShared::accept_bytes(&shared, 8);
2074 lock(&shared).stage.extend_from_slice(b"abcdefgh");
2075 let mut trailer = TrailerRecv::new(shared.clone());
2076 trailer.set_manual_credit();
2077
2078 let mut out = [0u8; 8];
2079 assert!(matches!(
2080 poll_read_once(&mut trailer, &mut out),
2081 Poll::Ready(Ok(8))
2082 ));
2083 assert!(lock(&sink.credits).is_empty(), "reading must not credit");
2084
2085 trailer.release(8);
2086 assert_eq!(&*lock(&sink.credits), &[(7, 8)]);
2087 }
2088
2089 #[test]
2092 fn credit_is_coalesced_below_half_an_interval() {
2093 let (shared, sink) = recv_shared(Limits {
2094 trailer_credit_interval: 64,
2095 trailer_session_window: 1024,
2096 ..Limits::default()
2097 });
2098 RecvShared::accept_bytes(&shared, 32);
2099 let mut trailer = TrailerRecv::new(shared.clone());
2100 trailer.set_manual_credit();
2101
2102 trailer.release(8);
2103 assert!(lock(&sink.credits).is_empty());
2104 trailer.release(8);
2105 assert!(lock(&sink.credits).is_empty());
2106 trailer.release(16);
2109 assert_eq!(&*lock(&sink.credits), &[(7, 32)]);
2110 }
2111
2112 #[test]
2116 fn credit_is_flushed_when_the_peer_is_starved() {
2117 let (shared, sink) = recv_shared(Limits {
2118 trailer_credit_interval: 64,
2119 trailer_session_window: 64,
2120 ..Limits::default()
2121 });
2122 assert_eq!(RecvShared::accept_bytes(&shared, 64), None);
2124 let mut trailer = TrailerRecv::new(shared.clone());
2125 trailer.set_manual_credit();
2126
2127 trailer.release(1);
2128 assert_eq!(
2129 &*lock(&sink.credits),
2130 &[(7, 1)],
2131 "a starved peer must be credited immediately, however little"
2132 );
2133 }
2134
2135 #[test]
2140 fn credit_is_flushed_when_the_consumer_starts_waiting() {
2141 let (shared, sink) = recv_shared(Limits {
2142 trailer_credit_interval: 64,
2143 trailer_session_window: 1024,
2144 ..Limits::default()
2145 });
2146 RecvShared::accept_bytes(&shared, 8);
2147 lock(&shared).stage.extend_from_slice(b"abcdefgh");
2148 let mut trailer = TrailerRecv::new(shared.clone());
2149
2150 let mut out = [0u8; 8];
2151 assert!(matches!(
2152 poll_read_once(&mut trailer, &mut out),
2153 Poll::Ready(Ok(8))
2154 ));
2155 assert!(
2156 lock(&sink.credits).is_empty(),
2157 "8 bytes is below half the interval, so it coalesces while reading"
2158 );
2159
2160 assert!(poll_read_once(&mut trailer, &mut out).is_pending());
2163 assert_eq!(&*lock(&sink.credits), &[(7, 8)]);
2164 }
2165
2166 #[test]
2170 fn credit_released_while_waiting_is_flushed_immediately() {
2171 let (shared, sink) = recv_shared(Limits {
2172 trailer_credit_interval: 64,
2173 trailer_session_window: 1024,
2174 ..Limits::default()
2175 });
2176 RecvShared::accept_bytes(&shared, 8);
2177 lock(&shared).stage.extend_from_slice(b"abcdefgh");
2178 let mut trailer = TrailerRecv::new(shared.clone());
2179 trailer.set_manual_credit();
2180
2181 let mut out = [0u8; 8];
2182 assert!(matches!(
2183 poll_read_once(&mut trailer, &mut out),
2184 Poll::Ready(Ok(8))
2185 ));
2186 assert!(poll_read_once(&mut trailer, &mut out).is_pending());
2187 assert!(
2188 lock(&sink.credits).is_empty(),
2189 "manual mode retires nothing on read, so the stall flushes nothing"
2190 );
2191
2192 trailer.release(8);
2193 assert_eq!(
2194 &*lock(&sink.credits),
2195 &[(7, 8)],
2196 "a waiting consumer's release goes out below the threshold"
2197 );
2198 }
2199
2200 #[test]
2206 fn credit_is_flushed_when_another_trailer_drained_the_pool() {
2207 let limits = Limits {
2208 trailer_credit_interval: 1024,
2209 trailer_session_window: 64,
2210 ..Limits::default()
2211 };
2212 let session = Arc::new(SessionWindow::new(limits.trailer_session_window));
2213 let make = |id| {
2214 let sink = Arc::new(RecordingSink::default());
2215 let shared = RecvShared::new(
2216 limits.trailer_recv_copy_threshold,
2217 limits.trailer_recv_demand_copy_threshold,
2218 limits.trailer_credit_interval,
2219 session.clone(),
2220 id,
2221 Arc::new(sink.clone()),
2222 );
2223 (shared, sink)
2224 };
2225 let (first, first_sink) = make(1);
2226 let (second, _second_sink) = make(2);
2227
2228 assert_eq!(RecvShared::accept_bytes(&first, 1), None);
2231 assert_eq!(RecvShared::accept_bytes(&second, 63), None);
2232
2233 let mut trailer = TrailerRecv::new(first.clone());
2234 trailer.set_manual_credit();
2235 trailer.release(1);
2236 assert_eq!(
2237 &*lock(&first_sink.credits),
2238 &[(1, 1)],
2239 "a drained pool must flush even a single byte, whoever drained it"
2240 );
2241 }
2242
2243 #[test]
2247 fn session_ledger_settles_each_byte_exactly_once() {
2248 let session = SessionWindow::new(100);
2249 assert_eq!(session.debit_up_to(1, 40), 40);
2250 assert_eq!(session.debit_up_to(2, 30), 30);
2251 assert_eq!(session.available(), 30);
2252
2253 session.refund(1, 1000);
2255 assert_eq!(session.available(), 70);
2256 session.refund(1, 1000);
2258 assert_eq!(session.available(), 70);
2259
2260 session.settle(2);
2262 assert_eq!(session.available(), 100);
2263 session.settle(2);
2264 session.refund(2, 30);
2265 assert_eq!(session.available(), 100);
2266
2267 session.refund(99, 10);
2270 session.settle(99);
2271 assert_eq!(session.available(), 100);
2272 }
2273
2274 #[test]
2278 fn reading_staged_bytes_after_a_discard_does_not_double_refund() {
2279 let limits = Limits {
2280 trailer_credit_interval: 8,
2281 trailer_session_window: 64,
2282 ..Limits::default()
2283 };
2284 let session = Arc::new(SessionWindow::new(limits.trailer_session_window));
2285 let sink = Arc::new(RecordingSink::default());
2286 let shared = RecvShared::new(
2287 limits.trailer_recv_copy_threshold,
2288 limits.trailer_recv_demand_copy_threshold,
2289 limits.trailer_credit_interval,
2290 session.clone(),
2291 7,
2292 Arc::new(sink.clone()),
2293 );
2294
2295 assert!(RecvShared::accept_bytes(&shared, 8).is_none());
2296 assert_eq!(session.available(), 56);
2297 lock(&shared).stage.extend_from_slice(b"abcdefgh");
2298
2299 let mut trailer = TrailerRecv::new(shared.clone());
2300 RecvShared::discard(&shared);
2301 assert_eq!(session.available(), 64, "the debt is returned in full");
2302
2303 let mut out = [0u8; 8];
2305 assert!(matches!(
2306 poll_read_once(&mut trailer, &mut out),
2307 Poll::Ready(Ok(8))
2308 ));
2309 assert_eq!(session.available(), 64);
2310 assert!(lock(&sink.credits).is_empty());
2311 }
2312
2313 #[test]
2316 fn dropping_a_trailer_discards_eagerly_exactly_once() {
2317 let (shared, sink) = recv_shared(unbounded_limits());
2318 let trailer = TrailerRecv::new(shared.clone());
2319 RecvShared::discard(&shared);
2320 assert_eq!(&*lock(&sink.discards), &[7]);
2321 drop(trailer);
2322 assert_eq!(&*lock(&sink.discards), &[7], "idempotent");
2323 }
2324}