1use std::time::Duration;
7
8use bytes::Bytes;
9use tokio::io::{AsyncReadExt, AsyncWrite, AsyncWriteExt};
10use tokio::net::TcpStream;
11use tokio::sync::{mpsc, watch};
12use tokio::task::JoinHandle;
13
14use microsandbox_protocol::bulk::{
15 BULK_FLOW_MASK_GUEST_TO_HOST, BULK_FLOW_MASK_HOST_TO_GUEST, BulkAccepted, BulkCredit,
16 BulkFinish, BulkFlow, BulkKind, BulkOffer, BulkReceiveState, BulkRecord, BulkSendState,
17 DEFAULT_BULK_RECORD_PAYLOAD, DEFAULT_BULK_WINDOW, MIN_BULK_RECORD_PAYLOAD,
18};
19use microsandbox_protocol::codec;
20use microsandbox_protocol::message::{Message, MessageType};
21use microsandbox_protocol::tcp::{TcpClosed, TcpConnect, TcpConnected, TcpData, TcpEof, TcpFailed};
22
23use crate::agent::{AdmittedBulkRecord, BulkInputPermit};
24use crate::serial::InputCharge;
25#[cfg(test)]
26use crate::session::SessionOutputEnvelope;
27use crate::session::{
28 BulkSessionOutput, RawActivity, RawSessionCompletion, RawSessionOutput, SessionOutput,
29 SessionOutputPermit, SessionOutputSender,
30};
31
32const TCP_CHUNK_SIZE: usize = 64 * 1024;
38
39const TCP_OUTPUT_RESERVATION: usize = 2 * TCP_CHUNK_SIZE;
44
45const TCP_COMMAND_CAPACITY: usize = DEFAULT_BULK_WINDOW as usize / MIN_BULK_RECORD_PAYLOAD as usize;
48
49const TCP_CONNECT_TIMEOUT: Duration = Duration::from_secs(30);
53
54pub struct TcpSession {
60 owner_id: u32,
61 commands: mpsc::Sender<TcpCommand>,
62 bulk_control: Option<TcpBulkControlSenders>,
63 task: JoinHandle<()>,
64 bulk: bool,
65}
66
67enum TcpCommand {
68 Data(Vec<u8>, Option<InputCharge>),
69 Eof(Option<InputCharge>),
70 BulkRecord(AdmittedBulkRecord),
71}
72
73struct TcpBulkControlSenders {
75 credit: watch::Sender<Option<BulkCredit>>,
76 finish: mpsc::Sender<BulkFinish>,
77}
78
79struct TcpBulkControlReceivers {
80 credit: watch::Receiver<Option<BulkCredit>>,
81 finish: mpsc::Receiver<BulkFinish>,
82}
83
84struct TcpBulkState {
85 send: BulkSendState,
86 receive: BulkReceiveState,
87}
88
89struct PendingTcpWrite {
92 payload: Bytes,
93 written: usize,
94 bulk_end: Option<u64>,
95 _bulk_input_permit: Option<BulkInputPermit>,
96 _control_input_charge: Option<InputCharge>,
97}
98
99impl TcpSession {
104 pub fn owner_id(&self) -> u32 {
106 self.owner_id
107 }
108
109 pub async fn write_data(&self, data: Vec<u8>) -> Result<(), String> {
114 self.write_data_charged(data, None).await
115 }
116
117 pub(crate) async fn write_data_charged(
118 &self,
119 data: Vec<u8>,
120 charge: Option<InputCharge>,
121 ) -> Result<(), String> {
122 if self.bulk {
123 return Err("CBOR TCP data is invalid after raw bulk acceptance".into());
124 }
125 self.commands
126 .send(TcpCommand::Data(data, charge))
127 .await
128 .map_err(|_| "TCP session is closed".to_string())
129 }
130
131 pub async fn close_write(&self) -> Result<(), String> {
136 self.close_write_charged(None).await
137 }
138
139 pub(crate) async fn close_write_charged(
140 &self,
141 charge: Option<InputCharge>,
142 ) -> Result<(), String> {
143 if self.bulk {
144 return Err("CBOR TCP EOF is invalid after raw bulk acceptance".into());
145 }
146 self.commands
147 .send(TcpCommand::Eof(charge))
148 .await
149 .map_err(|_| "TCP session is closed".to_string())
150 }
151
152 pub(crate) async fn write_bulk(&self, record: AdmittedBulkRecord) -> Result<(), String> {
154 if !self.bulk {
155 return Err("raw bulk record sent to a generation-6 TCP stream".into());
156 }
157 self.commands
158 .try_send(TcpCommand::BulkRecord(record))
159 .map_err(|error| format!("TCP bulk input queue is unavailable: {error}"))
160 }
161
162 pub async fn apply_credit(&self, credit: BulkCredit) -> Result<(), String> {
164 if !self.bulk {
165 return Err("bulk credit sent to a generation-6 TCP stream".into());
166 }
167 let control = self
168 .bulk_control
169 .as_ref()
170 .ok_or_else(|| "TCP bulk control path is unavailable".to_string())?;
171 if control.credit.is_closed() {
172 return Ok(());
175 }
176 control.credit.send_replace(Some(credit));
177 Ok(())
178 }
179
180 pub async fn finish_bulk(&self, finish: BulkFinish) -> Result<(), String> {
182 if !self.bulk {
183 return Err("bulk finish sent to a generation-6 TCP stream".into());
184 }
185 let control = self
186 .bulk_control
187 .as_ref()
188 .ok_or_else(|| "TCP bulk control path is unavailable".to_string())?;
189 control
190 .finish
191 .try_send(finish)
192 .map_err(|error| format!("TCP bulk finish path is unavailable: {error}"))
193 }
194
195 pub fn is_bulk(&self) -> bool {
197 self.bulk
198 }
199
200 pub fn close(&self) {
207 self.task.abort();
208 }
209
210 pub fn is_finished(&self) -> bool {
212 self.task.is_finished()
213 }
214
215 pub fn open(id: u32, req: TcpConnect, session_tx: &SessionOutputSender) -> Self {
224 let bulk = req.bulk.is_some();
225 let (commands_tx, commands_rx) = mpsc::channel(TCP_COMMAND_CAPACITY);
226 let (bulk_control, bulk_control_rx) = if bulk {
227 let (credit_tx, credit_rx) = watch::channel(None);
228 let (finish_tx, finish_rx) = mpsc::channel(1);
229 (
230 Some(TcpBulkControlSenders {
231 credit: credit_tx,
232 finish: finish_tx,
233 }),
234 Some(TcpBulkControlReceivers {
235 credit: credit_rx,
236 finish: finish_rx,
237 }),
238 )
239 } else {
240 (None, None)
241 };
242 let output_tx = session_tx.clone();
243 let task = tokio::spawn(async move {
244 connect_and_relay(id, req, commands_rx, bulk_control_rx, output_tx).await;
245 });
246
247 Self {
248 owner_id: id,
249 commands: commands_tx,
250 bulk_control,
251 task,
252 bulk,
253 }
254 }
255}
256
257async fn connect_and_relay(
268 id: u32,
269 req: TcpConnect,
270 commands: mpsc::Receiver<TcpCommand>,
271 bulk_control: Option<TcpBulkControlReceivers>,
272 tx: SessionOutputSender,
273) {
274 let TcpConnect { host, port, bulk } = req;
275 let connect = TcpStream::connect((host.as_str(), port));
276 let stream = match tokio::time::timeout(TCP_CONNECT_TIMEOUT, connect).await {
277 Ok(Ok(stream)) => stream,
278 Ok(Err(e)) => {
279 send_raw_tcp_message(
280 id,
281 MessageType::TcpFailed,
282 &TcpFailed {
283 error: format!("connect {host}:{port}: {e}"),
284 },
285 RawActivity::guest_message(),
286 Some(RawSessionCompletion::Tcp),
287 &tx,
288 )
289 .await;
290 return;
291 }
292 Err(_elapsed) => {
293 send_raw_tcp_message(
294 id,
295 MessageType::TcpFailed,
296 &TcpFailed {
297 error: format!("connect {host}:{port} timed out"),
298 },
299 RawActivity::guest_message(),
300 Some(RawSessionCompletion::Tcp),
301 &tx,
302 )
303 .await;
304 return;
305 }
306 };
307
308 if !send_raw_tcp_message(
309 id,
310 MessageType::TcpConnected,
311 &TcpConnected {},
312 RawActivity::guest_message(),
313 None,
314 &tx,
315 )
316 .await
317 {
318 return;
319 }
320
321 let bulk = match bulk {
322 Some(offer) => {
323 let accepted = match accept_tcp_offer(offer) {
324 Ok(accepted) => accepted,
325 Err(error) => {
326 send_raw_tcp_message(
327 id,
328 MessageType::TcpFailed,
329 &TcpFailed { error },
330 RawActivity::guest_message(),
331 Some(RawSessionCompletion::Tcp),
332 &tx,
333 )
334 .await;
335 return;
336 }
337 };
338 if !send_raw_tcp_message(
339 id,
340 MessageType::BulkAccepted,
341 &accepted,
342 RawActivity::guest_message(),
343 None,
344 &tx,
345 )
346 .await
347 {
348 return;
349 }
350 let send = match BulkSendState::new(
351 BulkKind::Tcp,
352 BulkFlow::GuestToHost,
353 accepted.max_record_payload,
354 accepted.guest_to_host_credit_limit,
355 ) {
356 Ok(send) => send,
357 Err(error) => {
358 eprintln!("failed to create TCP bulk send state for {id}: {error}");
359 return;
360 }
361 };
362 let receive = match BulkReceiveState::new(
363 BulkKind::Tcp,
364 BulkFlow::HostToGuest,
365 accepted.max_record_payload,
366 accepted.host_to_guest_credit_limit,
367 DEFAULT_BULK_WINDOW,
368 ) {
369 Ok(receive) => receive,
370 Err(error) => {
371 eprintln!("failed to create TCP bulk receive state for {id}: {error}");
372 return;
373 }
374 };
375 Some(TcpBulkState { send, receive })
376 }
377 None => None,
378 };
379
380 relay_tcp_session(id, stream, commands, bulk_control, tx, bulk).await;
381}
382
383fn accept_tcp_offer(offer: BulkOffer) -> Result<BulkAccepted, String> {
384 let offer = offer
385 .validate()
386 .map_err(|error| format!("invalid TCP bulk offer: {error}"))?;
387 if offer.guest_to_host_credit_limit == 0 {
388 return Err("TCP bulk offer must grant guest-to-host credit".into());
389 }
390 Ok(BulkAccepted {
391 kind: BulkKind::Tcp,
392 flows: BULK_FLOW_MASK_HOST_TO_GUEST | BULK_FLOW_MASK_GUEST_TO_HOST,
393 format: offer.format,
394 max_record_payload: offer.max_record_payload.min(DEFAULT_BULK_RECORD_PAYLOAD),
395 host_to_guest_credit_limit: DEFAULT_BULK_WINDOW,
396 guest_to_host_credit_limit: offer.guest_to_host_credit_limit,
397 })
398}
399
400async fn recv_optional_mpsc<T>(receiver: &mut Option<mpsc::Receiver<T>>) -> Option<T> {
402 loop {
403 let Some(active) = receiver.as_mut() else {
404 return std::future::pending().await;
405 };
406 match active.recv().await {
407 Some(value) => return Some(value),
408 None => *receiver = None,
409 }
410 }
411}
412
413async fn recv_optional_credit(
415 receiver: &mut Option<watch::Receiver<Option<BulkCredit>>>,
416) -> Option<BulkCredit> {
417 loop {
418 let Some(active) = receiver.as_mut() else {
419 return std::future::pending().await;
420 };
421 if active.changed().await.is_err() {
422 *receiver = None;
423 continue;
424 }
425 if let Some(credit) = *active.borrow_and_update() {
426 return Some(credit);
427 }
428 }
429}
430
431async fn apply_pending_tcp_finish<W>(
433 stream: &mut W,
434 state: &mut TcpBulkState,
435 pending: &mut Option<BulkFinish>,
436) -> Result<(), String>
437where
438 W: AsyncWrite + Unpin,
439{
440 let Some(finish) = *pending else {
441 return Ok(());
442 };
443 if finish.kind != BulkKind::Tcp || finish.flow != BulkFlow::HostToGuest {
444 return Err("bulk finish does not describe the TCP host-to-guest flow".into());
445 }
446 if finish.final_offset > state.receive.next_expected_offset() {
447 return Ok(());
448 }
449 state
450 .receive
451 .accept_finish(finish)
452 .map_err(|error| error.to_string())?;
453 stream
454 .shutdown()
455 .await
456 .map_err(|error| format!("shutdown TCP stream: {error}"))?;
457 *pending = None;
458 Ok(())
459}
460
461async fn relay_tcp_session(
462 id: u32,
463 stream: TcpStream,
464 mut commands: mpsc::Receiver<TcpCommand>,
465 bulk_control: Option<TcpBulkControlReceivers>,
466 tx: SessionOutputSender,
467 mut bulk: Option<TcpBulkState>,
468) {
469 let (mut reader, mut writer) = stream.into_split();
472 let (mut credit_rx, mut finish_rx) = match bulk_control {
473 Some(control) => (Some(control.credit), Some(control.finish)),
474 None => (None, None),
475 };
476 let mut pending_finish = None;
477 let read_capacity = bulk.as_ref().map_or(TCP_CHUNK_SIZE, |state| {
478 state.send.max_record_payload() as usize
479 });
480 let mut read_buf = vec![0u8; read_capacity];
481 let mut terminal_sent = false;
482 let mut pending_write: Option<PendingTcpWrite> = None;
483 let mut write_shutdown = false;
484 let mut read_eof = false;
487
488 loop {
489 if read_eof && write_shutdown {
493 break;
494 }
495 let read_limit = bulk.as_ref().map_or(TCP_CHUNK_SIZE, |state| {
496 state
497 .send
498 .available_credit()
499 .min(state.send.max_record_payload() as u64) as usize
500 });
501 tokio::select! {
502 Some(finish) = recv_optional_mpsc(&mut finish_rx) => {
503 if pending_finish.replace(finish).is_some() {
504 terminal_sent = send_tcp_failure(
505 id,
506 "duplicate TCP bulk finish".into(),
507 &tx,
508 )
509 .await;
510 break;
511 }
512 let Some(state) = bulk.as_mut() else {
513 terminal_sent = send_tcp_failure(
514 id,
515 "bulk finish received on a generation-6 TCP stream".into(),
516 &tx,
517 )
518 .await;
519 break;
520 };
521 if pending_write.is_none() {
522 if let Err(error) = apply_pending_tcp_finish(
523 &mut writer,
524 state,
525 &mut pending_finish,
526 ).await {
527 terminal_sent = send_tcp_failure(
528 id,
529 format!("invalid TCP bulk finish: {error}"),
530 &tx,
531 )
532 .await;
533 break;
534 }
535 write_shutdown = pending_finish.is_none();
536 }
537 }
538 Some(credit) = recv_optional_credit(&mut credit_rx) => {
539 let Some(state) = bulk.as_mut() else {
540 terminal_sent = send_tcp_failure(
541 id,
542 "bulk credit received on a generation-6 TCP stream".into(),
543 &tx,
544 )
545 .await;
546 break;
547 };
548 if let Err(error) = state.send.apply_credit(credit) {
549 terminal_sent = send_tcp_failure(
550 id,
551 format!("invalid TCP bulk credit: {error}"),
552 &tx,
553 )
554 .await;
555 break;
556 }
557 }
558 read = reader.read(&mut read_buf[..read_limit]), if !read_eof && read_limit != 0 => {
559 match read {
560 Ok(0) => {
561 if let Some(state) = bulk.as_mut() {
562 match state.send.finish() {
563 Ok(finish) => {
564 send_raw_tcp_message(
565 id,
566 MessageType::BulkFinish,
567 &finish,
568 RawActivity::guest_message(),
569 None,
570 &tx,
571 )
572 .await;
573 }
574 Err(error) => {
575 eprintln!("failed to finish TCP bulk receive flow {id}: {error}");
576 break;
577 }
578 }
579 } else {
580 send_raw_tcp_message(
581 id,
582 MessageType::TcpEof,
583 &TcpEof {},
584 RawActivity::guest_message(),
585 None,
586 &tx,
587 )
588 .await;
589 }
590 read_eof = true;
591 }
592 Ok(n) => {
593 if let Some(state) = bulk.as_mut() {
594 let offset = match state.send.admit(n) {
595 Ok(offset) => offset,
596 Err(error) => {
597 eprintln!("failed to admit TCP bulk record {id}: {error}");
598 break;
599 }
600 };
601 let Some(permit) = tx.reserve_bulk(n).await else {
602 break;
603 };
604 let record = BulkRecord {
605 id,
606 kind: BulkKind::Tcp,
607 flow: BulkFlow::GuestToHost,
608 offset,
609 payload: Bytes::copy_from_slice(&read_buf[..n]),
610 };
611 if !tx
612 .send_reserved(
613 id,
614 SessionOutput::Bulk(BulkSessionOutput::new(
615 record,
616 RawActivity::tcp_bytes(n),
617 )),
618 permit,
619 )
620 .await
621 {
622 break;
623 }
624 } else {
625 let Some(permit) = tx.reserve(TCP_OUTPUT_RESERVATION).await else {
626 break;
627 };
628 let data = read_buf[..n].to_vec();
629 if !send_raw_tcp_data(id, data, n, permit, &tx).await {
630 break;
631 }
632 }
633 }
634 Err(e) => {
635 terminal_sent = send_raw_tcp_message(
636 id,
637 MessageType::TcpFailed,
638 &TcpFailed {
639 error: format!("read TCP stream: {e}"),
640 },
641 RawActivity::guest_message(),
642 Some(RawSessionCompletion::Tcp),
643 &tx,
644 )
645 .await;
646 break;
647 }
648 }
649 }
650 write = async {
651 let pending = pending_write.as_ref().expect("guarded pending TCP write");
652 writer.write(&pending.payload[pending.written..]).await
653 }, if pending_write.is_some() => {
654 match write {
655 Ok(0) => {
656 terminal_sent = send_tcp_failure(
657 id,
658 "write TCP stream made no progress".into(),
659 &tx,
660 )
661 .await;
662 break;
663 }
664 Ok(written) => {
665 let pending = pending_write.as_mut().expect("guarded pending TCP write");
666 pending.written += written;
667 if pending.written != pending.payload.len() {
668 continue;
669 }
670
671 let completed = pending_write.take().expect("completed TCP write exists");
672 let bulk_end = completed.bulk_end;
673 drop(completed._bulk_input_permit);
676 drop(completed._control_input_charge);
677 if let Some(end) = bulk_end {
678 let Some(state) = bulk.as_mut() else {
679 terminal_sent = send_tcp_failure(
680 id,
681 "bulk TCP write lost its protocol state".into(),
682 &tx,
683 )
684 .await;
685 break;
686 };
687 match state.receive.consume(end) {
688 Ok(Some(credit)) => {
689 if !send_raw_tcp_message(
690 id,
691 MessageType::BulkCredit,
692 &credit,
693 RawActivity::guest_message(),
694 None,
695 &tx,
696 )
697 .await
698 {
699 break;
700 }
701 }
702 Ok(None) => {}
703 Err(error) => {
704 terminal_sent = send_tcp_failure(
705 id,
706 format!("advance TCP bulk credit: {error}"),
707 &tx,
708 )
709 .await;
710 break;
711 }
712 }
713 let finish_was_pending = pending_finish.is_some();
714 if let Err(error) = apply_pending_tcp_finish(
715 &mut writer,
716 state,
717 &mut pending_finish,
718 ).await {
719 terminal_sent = send_tcp_failure(
720 id,
721 format!("invalid TCP bulk finish: {error}"),
722 &tx,
723 )
724 .await;
725 break;
726 }
727 if finish_was_pending && pending_finish.is_none() {
728 write_shutdown = true;
729 }
730 }
731 }
732 Err(error) => {
733 terminal_sent = send_raw_tcp_message(
734 id,
735 MessageType::TcpFailed,
736 &TcpFailed {
737 error: format!("write TCP stream: {error}"),
738 },
739 RawActivity::guest_message(),
740 Some(RawSessionCompletion::Tcp),
741 &tx,
742 )
743 .await;
744 break;
745 }
746 }
747 }
748 command = commands.recv(), if pending_write.is_none() && !write_shutdown => {
749 match command {
750 Some(TcpCommand::Data(data, charge)) => {
751 if data.is_empty() {
752 drop(charge);
755 continue;
756 }
757 if bulk.is_some() {
758 terminal_sent = send_tcp_failure(
759 id,
760 "CBOR TCP data received after raw bulk acceptance".into(),
761 &tx,
762 )
763 .await;
764 break;
765 }
766 pending_write = Some(PendingTcpWrite {
767 payload: Bytes::from(data),
768 written: 0,
769 bulk_end: None,
770 _bulk_input_permit: None,
771 _control_input_charge: charge,
772 });
773 }
774 Some(TcpCommand::Eof(charge)) => {
775 if bulk.is_some() {
776 terminal_sent = send_tcp_failure(
777 id,
778 "CBOR TCP EOF received after raw bulk acceptance".into(),
779 &tx,
780 )
781 .await;
782 break;
783 }
784 if let Err(e) = writer.shutdown().await {
785 terminal_sent = send_raw_tcp_message(
786 id,
787 MessageType::TcpFailed,
788 &TcpFailed {
789 error: format!("shutdown TCP stream: {e}"),
790 },
791 RawActivity::guest_message(),
792 Some(RawSessionCompletion::Tcp),
793 &tx,
794 )
795 .await;
796 break;
797 }
798 write_shutdown = true;
799 drop(charge);
800 }
801 None => {
802 break;
803 }
804 Some(TcpCommand::BulkRecord(record)) => {
805 let Some(state) = bulk.as_mut() else {
806 terminal_sent = send_tcp_failure(
807 id,
808 "raw bulk record received on a generation-6 TCP stream".into(),
809 &tx,
810 )
811 .await;
812 break;
813 };
814 let end = match state.receive.accept_record(record.record()) {
815 Ok(end) => end,
816 Err(error) => {
817 terminal_sent = send_tcp_failure(
818 id,
819 format!("invalid TCP bulk record: {error}"),
820 &tx,
821 )
822 .await;
823 break;
824 }
825 };
826 let (record, permit) = record.into_parts();
827 pending_write = Some(PendingTcpWrite {
828 payload: record.payload,
829 written: 0,
830 bulk_end: Some(end),
831 _bulk_input_permit: Some(permit),
832 _control_input_charge: None,
833 });
834 }
835 }
836 }
837 }
838 }
839
840 if !terminal_sent {
841 send_raw_tcp_message(
842 id,
843 MessageType::TcpClosed,
844 &TcpClosed {},
845 RawActivity::guest_message(),
846 Some(RawSessionCompletion::Tcp),
847 &tx,
848 )
849 .await;
850 }
851}
852
853async fn send_tcp_failure(id: u32, error: String, tx: &SessionOutputSender) -> bool {
854 send_raw_tcp_message(
855 id,
856 MessageType::TcpFailed,
857 &TcpFailed { error },
858 RawActivity::guest_message(),
859 Some(RawSessionCompletion::Tcp),
860 tx,
861 )
862 .await
863}
864
865fn encode_tcp_message<T: serde::Serialize>(
866 id: u32,
867 t: MessageType,
868 payload: &T,
869 out_buf: &mut Vec<u8>,
870) -> Result<(), String> {
871 let msg = Message::with_payload(t, id, payload).map_err(|e| format!("encode tcp: {e}"))?;
872 codec::encode_to_buf(&msg, out_buf).map_err(|e| format!("encode tcp frame: {e}"))?;
873 Ok(())
874}
875
876async fn send_raw_tcp_message<T: serde::Serialize>(
877 id: u32,
878 t: MessageType,
879 payload: &T,
880 activity: RawActivity,
881 completion: Option<RawSessionCompletion>,
882 tx: &SessionOutputSender,
883) -> bool {
884 let mut buf = Vec::new();
885 match encode_tcp_message(id, t, payload, &mut buf) {
886 Ok(()) => {
887 tx.send(
888 id,
889 SessionOutput::Raw(RawSessionOutput::new(buf, activity, completion)),
890 )
891 .await
892 }
893 Err(e) => {
894 eprintln!("failed to encode tcp message for {id}: {e}");
895 false
896 }
897 }
898}
899
900async fn send_raw_tcp_data(
902 id: u32,
903 data: Vec<u8>,
904 byte_count: usize,
905 permit: SessionOutputPermit,
906 tx: &SessionOutputSender,
907) -> bool {
908 let mut buf = Vec::new();
909 match encode_tcp_message(id, MessageType::TcpData, &TcpData { data }, &mut buf) {
910 Ok(()) => {
911 tx.send_reserved(
912 id,
913 SessionOutput::Raw(RawSessionOutput::new(
914 buf,
915 RawActivity::tcp_bytes(byte_count),
916 None,
917 )),
918 permit,
919 )
920 .await
921 }
922 Err(error) => {
923 eprintln!("failed to encode TCP data for {id}: {error}");
924 false
925 }
926 }
927}
928
929#[cfg(test)]
934mod tests {
935 use std::time::Duration;
936
937 use microsandbox_protocol::message::FLAG_TERMINAL;
938 use tokio::net::TcpListener;
939
940 use super::*;
941
942 #[test]
943 fn admitted_transport_window_fits_each_tcp_input_queue() {
944 use microsandbox_protocol::core::{
945 WORKLOAD_TRANSPORT_BULK_FRAMES, WORKLOAD_TRANSPORT_CONTROL_FRAMES,
946 };
947
948 assert!(
951 WORKLOAD_TRANSPORT_CONTROL_FRAMES + WORKLOAD_TRANSPORT_BULK_FRAMES
952 <= TCP_COMMAND_CAPACITY as u64
953 );
954 }
955
956 #[tokio::test]
957 async fn blocked_tcp_retains_data_and_eof_credit_until_consumption_or_cancel() {
958 use crate::serial::{InputLane, InputWindow};
959 use microsandbox_protocol::core::WorkloadTransportCredit;
960 use std::os::fd::AsRawFd;
961
962 for cancel in [false, true] {
963 let listener = TcpListener::bind(("127.0.0.1", 0)).await.unwrap();
964 let receive_bytes: libc::c_int = 64 * 1024;
967 assert_eq!(
968 unsafe {
969 libc::setsockopt(
970 listener.as_raw_fd(),
971 libc::SOL_SOCKET,
972 libc::SO_RCVBUF,
973 (&receive_bytes as *const libc::c_int).cast(),
974 std::mem::size_of_val(&receive_bytes) as libc::socklen_t,
975 )
976 },
977 0
978 );
979 let (sender, mut output) = SessionOutputSender::channel();
980 let session = TcpSession::open(
981 8,
982 TcpConnect {
983 host: "127.0.0.1".into(),
984 port: listener.local_addr().unwrap().port(),
985 bulk: None,
986 },
987 &sender,
988 );
989 let (mut peer, _) = listener.accept().await.unwrap();
990 assert_eq!(recv_message(&mut output).await.t, MessageType::TcpConnected);
991 let initial = WorkloadTransportCredit {
992 control_bytes: 64,
993 control_frames: 2,
994 bulk_bytes: 8 * 1024 * 1024,
995 bulk_frames: 2,
996 };
997 let ledger = InputWindow::new(initial);
998 let payload_len = initial.bulk_bytes as usize - 64;
999 let data_charge = ledger.admit(InputLane::Bulk, payload_len + 32).unwrap();
1000 let eof_charge = ledger.admit(InputLane::Bulk, 32).unwrap();
1001 tokio::time::timeout(Duration::from_millis(100), async {
1002 session
1003 .write_data_charged(vec![0x5c; payload_len], Some(data_charge))
1004 .await
1005 .unwrap();
1006 session.close_write_charged(Some(eof_charge)).await.unwrap();
1007 })
1008 .await
1009 .expect("admitted input waited for a blocked TCP consumer");
1010 tokio::time::sleep(Duration::from_millis(20)).await;
1011 assert_eq!(ledger.credit().unwrap(), initial);
1012 assert!(ledger.admit(InputLane::Bulk, 1).is_err());
1013 if cancel {
1014 session.close();
1015 wait_finished(&session).await;
1016 } else {
1017 let mut bytes = Vec::new();
1018 tokio::time::timeout(Duration::from_secs(10), peer.read_to_end(&mut bytes))
1019 .await
1020 .expect("ordered TCP EOF did not arrive")
1021 .unwrap();
1022 assert_eq!(bytes.len(), payload_len);
1023 assert!(bytes.iter().all(|byte| *byte == 0x5c));
1024 session.close();
1025 wait_finished(&session).await;
1026 }
1027 assert_eq!(ledger.credit().unwrap().bulk_bytes, initial.bulk_bytes * 2);
1028 assert_eq!(ledger.credit().unwrap().bulk_frames, 4);
1029 assert_eq!(
1030 ledger.credit().unwrap().control_bytes,
1031 initial.control_bytes
1032 );
1033 }
1034 }
1035
1036 #[tokio::test]
1037 async fn connect_failure_sends_terminal_failed() {
1038 let (session_tx, mut session_rx) = SessionOutputSender::channel();
1039
1040 let session = TcpSession::open(
1041 7,
1042 TcpConnect {
1043 host: "127.0.0.1".to_string(),
1044 port: 0,
1045 bulk: None,
1046 },
1047 &session_tx,
1048 );
1049
1050 let msg = recv_message(&mut session_rx).await;
1052 assert_eq!(msg.t, MessageType::TcpFailed);
1053 assert_eq!(msg.flags, FLAG_TERMINAL);
1054 let failed: TcpFailed = msg.payload().unwrap();
1055 assert!(failed.error.contains("connect 127.0.0.1:0"));
1056
1057 wait_finished(&session).await;
1058 }
1059
1060 #[tokio::test]
1061 async fn close_request_finishes_session_task() {
1062 let listener = TcpListener::bind(("127.0.0.1", 0)).await.unwrap();
1063 let port = listener.local_addr().unwrap().port();
1064 let (session_tx, mut session_rx) = SessionOutputSender::channel();
1065 let accept_task = tokio::spawn(async move {
1066 let (_socket, _) = listener.accept().await.unwrap();
1067 tokio::time::sleep(Duration::from_secs(5)).await;
1068 });
1069
1070 let session = TcpSession::open(
1071 9,
1072 TcpConnect {
1073 host: "127.0.0.1".to_string(),
1074 port,
1075 bulk: None,
1076 },
1077 &session_tx,
1078 );
1079
1080 let connected = recv_message(&mut session_rx).await;
1081 assert_eq!(connected.t, MessageType::TcpConnected);
1082
1083 session.close();
1084 wait_finished(&session).await;
1085
1086 accept_task.abort();
1087 }
1088
1089 #[tokio::test]
1090 async fn destination_eof_keeps_session_open_for_host_writes() {
1091 let listener = TcpListener::bind(("127.0.0.1", 0)).await.unwrap();
1092 let port = listener.local_addr().unwrap().port();
1093 let (session_tx, mut session_rx) = SessionOutputSender::channel();
1094
1095 let (got_tx, got_rx) = tokio::sync::oneshot::channel();
1098 let accept_task = tokio::spawn(async move {
1099 let (mut socket, _) = listener.accept().await.unwrap();
1100 socket.shutdown().await.unwrap();
1101 let mut buf = Vec::new();
1102 socket.read_to_end(&mut buf).await.unwrap();
1103 let _ = got_tx.send(buf);
1104 });
1105
1106 let session = TcpSession::open(
1107 11,
1108 TcpConnect {
1109 host: "127.0.0.1".to_string(),
1110 port,
1111 bulk: None,
1112 },
1113 &session_tx,
1114 );
1115
1116 let connected = recv_message(&mut session_rx).await;
1117 assert_eq!(connected.t, MessageType::TcpConnected);
1118
1119 let eof = recv_message(&mut session_rx).await;
1122 assert_eq!(eof.t, MessageType::TcpEof);
1123 assert_ne!(eof.flags, FLAG_TERMINAL);
1124 assert!(!session.is_finished());
1125
1126 session.write_data(b"after-eof".to_vec()).await.unwrap();
1128 session.close_write().await.unwrap();
1129 let received = tokio::time::timeout(Duration::from_secs(1), got_rx)
1130 .await
1131 .unwrap()
1132 .unwrap();
1133 assert_eq!(received, b"after-eof");
1134
1135 session.close();
1137 wait_finished(&session).await;
1138
1139 accept_task.await.unwrap();
1140 }
1141
1142 #[tokio::test]
1143 async fn active_raw_credit_validation_and_inline_negotiation_still_apply() {
1144 for raw in [false, true] {
1145 let listener = TcpListener::bind(("127.0.0.1", 0)).await.unwrap();
1146 let (tx, mut rx) = SessionOutputSender::channel();
1147 let session = TcpSession::open(
1148 41,
1149 TcpConnect {
1150 host: "127.0.0.1".into(),
1151 port: listener.local_addr().unwrap().port(),
1152 bulk: raw.then(BulkOffer::tcp),
1153 },
1154 &tx,
1155 );
1156 let (_peer, _) = listener.accept().await.unwrap();
1157 assert_eq!(recv_message(&mut rx).await.t, MessageType::TcpConnected);
1158 if raw {
1159 assert_eq!(recv_message(&mut rx).await.t, MessageType::BulkAccepted);
1160 }
1161 let result = session
1162 .apply_credit(BulkCredit {
1163 kind: BulkKind::Tcp,
1164 flow: BulkFlow::GuestToHost,
1165 consumed_offset: 1,
1166 credit_limit: DEFAULT_BULK_WINDOW + 1,
1167 })
1168 .await;
1169 if raw {
1170 result.unwrap();
1171 let failed = tokio::time::timeout(Duration::from_secs(1), recv_message(&mut rx))
1172 .await
1173 .unwrap();
1174 assert_eq!(failed.t, MessageType::TcpFailed);
1175 assert_eq!(failed.flags, FLAG_TERMINAL);
1176 assert!(
1177 failed
1178 .payload::<TcpFailed>()
1179 .unwrap()
1180 .error
1181 .contains("not admitted")
1182 );
1183 wait_finished(&session).await;
1184 } else {
1185 assert!(result.unwrap_err().contains("generation-6"));
1186 session.close();
1187 wait_finished(&session).await;
1188 }
1189 }
1190 }
1191
1192 #[tokio::test]
1193 async fn both_half_close_orders_preserve_data_and_emit_one_terminal() {
1194 for raw in [false, true] {
1195 for peer_first in [false, true] {
1196 tokio::time::timeout(Duration::from_secs(5), async {
1197 let listener = TcpListener::bind(("127.0.0.1", 0)).await.unwrap();
1198 let (tx, mut rx) = SessionOutputSender::channel();
1199 let session = TcpSession::open(
1200 31,
1201 TcpConnect {
1202 host: "127.0.0.1".into(),
1203 port: listener.local_addr().unwrap().port(),
1204 bulk: raw.then(BulkOffer::tcp),
1205 },
1206 &tx,
1207 );
1208 let (mut peer, _) = listener.accept().await.unwrap();
1209 assert_eq!(recv_message(&mut rx).await.t, MessageType::TcpConnected);
1210 if raw {
1211 assert_eq!(recv_message(&mut rx).await.t, MessageType::BulkAccepted);
1212 }
1213 let host_data = b"host data survives the peer's first EOF";
1214 let peer_data = b"peer data survives the host's first EOF";
1215 if peer_first {
1216 peer.write_all(peer_data).await.unwrap();
1217 peer.shutdown().await.unwrap();
1218 assert_tcp_output_through_eof(&mut rx, raw, peer_data).await;
1219 assert!(!session.is_finished(), "one EOF must preserve host writes");
1220 send_test_input_and_eof(&session, raw, host_data).await;
1221 } else {
1222 send_test_input_and_eof(&session, raw, host_data).await;
1223 }
1224
1225 let mut received = Vec::new();
1226 peer.read_to_end(&mut received).await.unwrap();
1227 assert_eq!(received, host_data);
1228 if !peer_first {
1229 assert!(!session.is_finished(), "one EOF must preserve peer output");
1230 peer.write_all(peer_data).await.unwrap();
1231 peer.shutdown().await.unwrap();
1232 assert_tcp_output_through_eof(&mut rx, raw, peer_data).await;
1233 }
1234 assert_one_normal_terminal(&session, &mut rx).await;
1235 })
1236 .await
1237 .unwrap_or_else(|_| {
1238 panic!("TCP completion timed out: raw={raw}, peer_first={peer_first}")
1239 });
1240 }
1241 }
1242 }
1243
1244 #[tokio::test]
1245 async fn raw_finish_waits_for_delayed_record_and_pending_socket_write_before_terminal() {
1246 use std::os::fd::AsRawFd;
1247
1248 tokio::time::timeout(Duration::from_secs(5), async {
1249 let listener = TcpListener::bind(("127.0.0.1", 0)).await.unwrap();
1250 let (stream, accepted) = tokio::join!(
1251 TcpStream::connect(listener.local_addr().unwrap()),
1252 listener.accept(),
1253 );
1254 let stream = stream.unwrap();
1255 let (mut peer, _) = accepted.unwrap();
1256 for (fd, option, bytes) in [
1259 (stream.as_raw_fd(), libc::SO_SNDBUF, 4096 as libc::c_int),
1260 (peer.as_raw_fd(), libc::SO_RCVBUF, 65536 as libc::c_int),
1261 ] {
1262 assert_eq!(
1263 unsafe {
1264 libc::setsockopt(
1265 fd,
1266 libc::SOL_SOCKET,
1267 option,
1268 (&bytes as *const libc::c_int).cast(),
1269 std::mem::size_of_val(&bytes) as libc::socklen_t,
1270 )
1271 },
1272 0
1273 );
1274 }
1275 let (tx, mut rx) = SessionOutputSender::channel();
1276 let (commands, commands_rx) = mpsc::channel(TCP_COMMAND_CAPACITY);
1277 let (credit, credit_rx) = watch::channel(None);
1278 let (finish, finish_rx) = mpsc::channel(1);
1279 let task = tokio::spawn(relay_tcp_session(
1280 37,
1281 stream,
1282 commands_rx,
1283 Some(TcpBulkControlReceivers {
1284 credit: credit_rx,
1285 finish: finish_rx,
1286 }),
1287 tx,
1288 Some(TcpBulkState {
1289 send: BulkSendState::new(
1290 BulkKind::Tcp,
1291 BulkFlow::GuestToHost,
1292 DEFAULT_BULK_RECORD_PAYLOAD,
1293 DEFAULT_BULK_WINDOW,
1294 )
1295 .unwrap(),
1296 receive: BulkReceiveState::new(
1297 BulkKind::Tcp,
1298 BulkFlow::HostToGuest,
1299 DEFAULT_BULK_RECORD_PAYLOAD,
1300 DEFAULT_BULK_WINDOW,
1301 DEFAULT_BULK_WINDOW,
1302 )
1303 .unwrap(),
1304 }),
1305 ));
1306 let session = TcpSession {
1307 owner_id: 37,
1308 commands,
1309 bulk_control: Some(TcpBulkControlSenders { credit, finish }),
1310 task,
1311 bulk: true,
1312 };
1313 peer.shutdown().await.unwrap();
1314 assert_tcp_output_through_eof(&mut rx, true, b"").await;
1315 let payload = Bytes::from(vec![0x6a; DEFAULT_BULK_RECORD_PAYLOAD as usize]);
1316 session
1317 .finish_bulk(BulkFinish {
1318 kind: BulkKind::Tcp,
1319 flow: BulkFlow::HostToGuest,
1320 final_offset: payload.len() as u64,
1321 })
1322 .await
1323 .unwrap();
1324 while session.bulk_control.as_ref().unwrap().finish.capacity() == 0 {
1325 tokio::task::yield_now().await;
1326 }
1327 assert!(
1328 !session.is_finished(),
1329 "finish cannot skip its missing final record"
1330 );
1331 assert!(matches!(
1332 rx.try_recv(),
1333 Err(mpsc::error::TryRecvError::Empty)
1334 ));
1335 session
1336 .write_bulk(AdmittedBulkRecord::for_test(BulkRecord {
1337 id: 37,
1338 kind: BulkKind::Tcp,
1339 flow: BulkFlow::HostToGuest,
1340 offset: 0,
1341 payload: payload.clone(),
1342 }))
1343 .await
1344 .unwrap();
1345 let mut received = vec![0];
1346 peer.read_exact(&mut received).await.unwrap();
1347 assert!(
1348 !session.is_finished(),
1349 "finish cannot skip a partial socket write"
1350 );
1351 assert!(matches!(
1352 rx.try_recv(),
1353 Err(mpsc::error::TryRecvError::Empty)
1354 ));
1355 peer.read_to_end(&mut received).await.unwrap();
1356 assert_eq!(received, payload);
1357 assert_one_normal_terminal(&session, &mut rx).await;
1358 })
1359 .await
1360 .expect("delayed raw record did not finish normally");
1361 }
1362
1363 #[tokio::test]
1364 async fn raw_bulk_tcp_relays_both_directions_and_exact_half_closes() {
1365 let listener = TcpListener::bind(("127.0.0.1", 0)).await.unwrap();
1366 let port = listener.local_addr().unwrap().port();
1367 let (session_tx, mut session_rx) = SessionOutputSender::channel();
1368 let (got_tx, got_rx) = tokio::sync::oneshot::channel();
1369 let accept_task = tokio::spawn(async move {
1370 let (mut socket, _) = listener.accept().await.unwrap();
1371 socket.write_all(b"from-destination").await.unwrap();
1372 socket.shutdown().await.unwrap();
1373 let mut received = Vec::new();
1374 socket.read_to_end(&mut received).await.unwrap();
1375 got_tx.send(received).unwrap();
1376 });
1377
1378 let session = TcpSession::open(
1379 13,
1380 TcpConnect {
1381 host: "127.0.0.1".to_string(),
1382 port,
1383 bulk: Some(BulkOffer::tcp()),
1384 },
1385 &session_tx,
1386 );
1387 assert_eq!(
1388 recv_message(&mut session_rx).await.t,
1389 MessageType::TcpConnected
1390 );
1391 assert_eq!(
1392 recv_message(&mut session_rx).await.t,
1393 MessageType::BulkAccepted
1394 );
1395
1396 let host_payload = Bytes::from_static(b"from-host");
1397 session
1398 .write_bulk(AdmittedBulkRecord::for_test(BulkRecord {
1399 id: 13,
1400 kind: BulkKind::Tcp,
1401 flow: BulkFlow::HostToGuest,
1402 offset: 0,
1403 payload: host_payload.clone(),
1404 }))
1405 .await
1406 .unwrap();
1407 session
1408 .finish_bulk(BulkFinish {
1409 kind: BulkKind::Tcp,
1410 flow: BulkFlow::HostToGuest,
1411 final_offset: host_payload.len() as u64,
1412 })
1413 .await
1414 .unwrap();
1415
1416 let record = recv_bulk(&mut session_rx).await;
1417 assert_eq!(record.flow, BulkFlow::GuestToHost);
1418 assert_eq!(record.offset, 0);
1419 assert_eq!(record.payload, Bytes::from_static(b"from-destination"));
1420 let finish = recv_message(&mut session_rx).await;
1421 assert_eq!(finish.t, MessageType::BulkFinish);
1422 let finish: BulkFinish = finish.payload().unwrap();
1423 assert_eq!(finish.final_offset, b"from-destination".len() as u64);
1424
1425 let received = tokio::time::timeout(Duration::from_secs(1), got_rx)
1426 .await
1427 .unwrap()
1428 .unwrap();
1429 assert_eq!(received, host_payload);
1430 session.close();
1431 wait_finished(&session).await;
1432 accept_task.await.unwrap();
1433 }
1434
1435 #[tokio::test]
1436 async fn blocked_host_to_guest_write_does_not_stop_guest_to_host_reads() {
1437 let listener = TcpListener::bind(("127.0.0.1", 0)).await.unwrap();
1438 let port = listener.local_addr().unwrap().port();
1439 let (session_tx, mut session_rx) = SessionOutputSender::channel();
1440 let host_payload = Bytes::from(vec![0x3c; DEFAULT_BULK_RECORD_PAYLOAD as usize]);
1441 let guest_payload = vec![0x7a; DEFAULT_BULK_WINDOW as usize];
1442 let expected_guest_payload = guest_payload.clone();
1443 let expected_host_payload = host_payload.clone();
1444 let host_payload_len = host_payload.len();
1445 let (got_tx, got_rx) = tokio::sync::oneshot::channel();
1446
1447 let accept_task = tokio::spawn(async move {
1451 let (mut socket, _) = listener.accept().await.unwrap();
1452 let mut received = vec![0u8; host_payload_len];
1453 socket.read_exact(&mut received[..1]).await.unwrap();
1454 socket.write_all(&guest_payload).await.unwrap();
1455 socket.read_exact(&mut received[1..]).await.unwrap();
1456 got_tx.send(received).unwrap();
1457 });
1458
1459 let session = TcpSession::open(
1460 17,
1461 TcpConnect {
1462 host: "127.0.0.1".to_string(),
1463 port,
1464 bulk: Some(BulkOffer::tcp()),
1465 },
1466 &session_tx,
1467 );
1468 assert_eq!(
1469 recv_message(&mut session_rx).await.t,
1470 MessageType::TcpConnected
1471 );
1472 assert_eq!(
1473 recv_message(&mut session_rx).await.t,
1474 MessageType::BulkAccepted
1475 );
1476
1477 session
1478 .write_bulk(AdmittedBulkRecord::for_test(BulkRecord {
1479 id: 17,
1480 kind: BulkKind::Tcp,
1481 flow: BulkFlow::HostToGuest,
1482 offset: 0,
1483 payload: host_payload,
1484 }))
1485 .await
1486 .unwrap();
1487
1488 let received_guest_payload = tokio::time::timeout(Duration::from_secs(5), async {
1489 let mut received = Vec::with_capacity(expected_guest_payload.len());
1490 while received.len() < expected_guest_payload.len() {
1491 let record = recv_bulk(&mut session_rx).await;
1492 assert_eq!(record.offset, received.len() as u64);
1493 received.extend_from_slice(&record.payload);
1494 }
1495 received
1496 })
1497 .await
1498 .expect("guest-to-host reads must progress while the opposite write is blocked");
1499 assert_eq!(received_guest_payload, expected_guest_payload);
1500
1501 let received_host_payload = tokio::time::timeout(Duration::from_secs(5), got_rx)
1502 .await
1503 .unwrap()
1504 .unwrap();
1505 assert_eq!(received_host_payload, expected_host_payload);
1506
1507 session.close();
1508 wait_finished(&session).await;
1509 accept_task.await.unwrap();
1510 }
1511
1512 #[tokio::test]
1513 async fn full_bulk_data_queue_does_not_starve_credit_or_finish() {
1514 let (commands, _commands_rx) = mpsc::channel(TCP_COMMAND_CAPACITY);
1515 let (credit, mut credit_rx) = watch::channel(None);
1516 let (finish, mut finish_rx) = mpsc::channel(1);
1517 let task = tokio::spawn(std::future::pending());
1518 let session = TcpSession {
1519 owner_id: 29,
1520 commands,
1521 bulk_control: Some(TcpBulkControlSenders { credit, finish }),
1522 task,
1523 bulk: true,
1524 };
1525 let payload = Bytes::from(vec![0u8; MIN_BULK_RECORD_PAYLOAD as usize]);
1526
1527 for index in 0..TCP_COMMAND_CAPACITY {
1528 session
1529 .write_bulk(AdmittedBulkRecord::for_test(BulkRecord {
1530 id: 29,
1531 kind: BulkKind::Tcp,
1532 flow: BulkFlow::HostToGuest,
1533 offset: (index * MIN_BULK_RECORD_PAYLOAD as usize) as u64,
1534 payload: payload.clone(),
1535 }))
1536 .await
1537 .unwrap();
1538 }
1539 assert_eq!(session.commands.capacity(), 0);
1540
1541 let credit_update = BulkCredit {
1542 kind: BulkKind::Tcp,
1543 flow: BulkFlow::GuestToHost,
1544 consumed_offset: 4,
1545 credit_limit: 8,
1546 };
1547 session.apply_credit(credit_update).await.unwrap();
1548 credit_rx.changed().await.unwrap();
1549 assert_eq!(*credit_rx.borrow_and_update(), Some(credit_update));
1550
1551 let finish_update = BulkFinish {
1552 kind: BulkKind::Tcp,
1553 flow: BulkFlow::HostToGuest,
1554 final_offset: DEFAULT_BULK_WINDOW,
1555 };
1556 session.finish_bulk(finish_update).await.unwrap();
1557 assert_eq!(finish_rx.recv().await, Some(finish_update));
1558
1559 session.close();
1560 }
1561
1562 async fn send_test_input_and_eof(session: &TcpSession, raw: bool, data: &[u8]) {
1563 if raw {
1564 session
1565 .write_bulk(AdmittedBulkRecord::for_test(BulkRecord {
1566 id: session.owner_id(),
1567 kind: BulkKind::Tcp,
1568 flow: BulkFlow::HostToGuest,
1569 offset: 0,
1570 payload: Bytes::copy_from_slice(data),
1571 }))
1572 .await
1573 .unwrap();
1574 session
1575 .finish_bulk(BulkFinish {
1576 kind: BulkKind::Tcp,
1577 flow: BulkFlow::HostToGuest,
1578 final_offset: data.len() as u64,
1579 })
1580 .await
1581 .unwrap();
1582 } else {
1583 session.write_data(data.to_vec()).await.unwrap();
1584 session.close_write().await.unwrap();
1585 }
1586 }
1587
1588 async fn assert_tcp_output_through_eof(
1589 rx: &mut mpsc::Receiver<SessionOutputEnvelope>,
1590 raw: bool,
1591 expected: &[u8],
1592 ) {
1593 let mut received = Vec::new();
1594 loop {
1595 let envelope = rx.recv().await.expect("TCP output ended before EOF");
1596 match envelope.output {
1597 SessionOutput::Bulk(output) => {
1598 assert!(raw);
1599 assert_eq!(output.record.offset, received.len() as u64);
1600 received.extend_from_slice(&output.record.payload);
1601 }
1602 SessionOutput::Raw(mut output) => {
1603 let message = decode_one_message(&mut output.frame);
1604 assert_eq!(message.flags & FLAG_TERMINAL, 0, "terminal preceded EOF");
1605 match message.t {
1606 MessageType::TcpData => {
1607 assert!(!raw);
1608 received.extend(message.payload::<TcpData>().unwrap().data);
1609 }
1610 MessageType::TcpEof => {
1611 assert!(!raw);
1612 break;
1613 }
1614 MessageType::BulkFinish => {
1615 assert!(raw);
1616 let finish = message.payload::<BulkFinish>().unwrap();
1617 assert_eq!(finish.final_offset, received.len() as u64);
1618 break;
1619 }
1620 _ => panic!("unexpected TCP output: {:?}", message.t),
1621 }
1622 }
1623 _ => panic!("unexpected non-TCP output"),
1624 }
1625 }
1626 assert_eq!(received, expected);
1627 }
1628
1629 async fn assert_one_normal_terminal(
1630 session: &TcpSession,
1631 rx: &mut mpsc::Receiver<SessionOutputEnvelope>,
1632 ) {
1633 let closed = recv_message(rx).await;
1634 assert_eq!(closed.t, MessageType::TcpClosed);
1635 assert_eq!(closed.flags, FLAG_TERMINAL);
1636 closed.payload::<TcpClosed>().unwrap();
1637 wait_finished(session).await;
1638 assert!(matches!(
1639 rx.try_recv(),
1640 Err(mpsc::error::TryRecvError::Empty | mpsc::error::TryRecvError::Disconnected)
1641 ));
1642 }
1643
1644 async fn wait_finished(session: &TcpSession) {
1645 tokio::time::timeout(Duration::from_secs(1), async {
1646 while !session.is_finished() {
1647 tokio::time::sleep(Duration::from_millis(10)).await;
1648 }
1649 })
1650 .await
1651 .unwrap();
1652 }
1653
1654 fn decode_one_message(buf: &mut Vec<u8>) -> Message {
1655 codec::try_decode_from_buf(buf).unwrap().unwrap()
1656 }
1657
1658 async fn recv_message(rx: &mut mpsc::Receiver<SessionOutputEnvelope>) -> Message {
1659 let envelope = rx.recv().await.unwrap();
1660 let SessionOutput::Raw(mut output) = envelope.output else {
1661 panic!("expected SessionOutput::Raw frame");
1662 };
1663 decode_one_message(&mut output.frame)
1664 }
1665
1666 async fn recv_bulk(rx: &mut mpsc::Receiver<SessionOutputEnvelope>) -> BulkRecord {
1667 let envelope = rx.recv().await.unwrap();
1668 let SessionOutput::Bulk(output) = envelope.output else {
1669 panic!("expected SessionOutput::Bulk record");
1670 };
1671 output.record
1672 }
1673}