1use super::envelope::{Header, IncomingEnvelope, MessageKind, Parity, Side};
7use super::operation::{
8 OperationHandle, OperationKey, OutgoingBody, OutgoingMessage, PendingOperation,
9};
10use super::promise::PromiseResult;
11use super::worker;
12use super::{
13 Closer, DEFAULT_AUTOREPLY_TIMEOUT, DEFAULT_MAX_INBOUND_BYTES, DEFAULT_MAX_INBOUND_REQUESTS,
14 Error, Message, Promise, Requester, Responder, schema,
15};
16use crate::LogId;
17use crate::transport::{self, Read, Stream, Verifier, Write};
18use prost::bytes::Bytes;
19use std::collections::{HashMap, HashSet, VecDeque};
20use std::fmt;
21use std::sync::atomic::{AtomicUsize, Ordering};
22use std::sync::{Arc, Condvar, Mutex};
23use std::time::{Duration, Instant};
24
25pub fn connect<R, W, V>(stream: Stream<R, W>, verifier: &V) -> Result<(Session, V::Info), Error>
31where
32 R: Read + Send + 'static,
33 W: Write + Send + 'static,
34 V: Verifier,
35{
36 let mut client = transport::Client::new(stream);
37
38 let (sender, info) = client.connect(verifier)?;
39 #[cfg(any(test, feature = "fuzz"))]
40 let workers = Arc::new(worker::Tracker::default());
41 let session = Session::start(
42 Side::Client,
43 sender,
44 Some(client.closer()),
45 #[cfg(any(test, feature = "fuzz"))]
46 workers.clone(),
47 );
48 let inner = session.inner.clone();
49
50 worker::spawn(
51 "wire-client-reader",
52 #[cfg(any(test, feature = "fuzz"))]
53 &workers,
54 move || run_reader(client, inner),
55 );
56 Ok((session, info))
57}
58
59fn run_reader<R: Read, W: Write>(mut client: transport::Client<R, W>, session: Arc<SessionInner>) {
63 loop {
64 let result = client
65 .recv()
66 .map_err(Error::from)
67 .and_then(|bytes| session.handle_message(bytes));
68 if let Err(error) = result {
69 session.close(error);
72 break;
73 }
74 }
75}
76
77pub struct Session {
97 pub(super) inner: Arc<SessionInner>,
99}
100
101impl Session {
102 pub fn set_autoreply_timeout(self, timeout: Duration) -> Self {
114 self.inner.set_autoreply_timeout(timeout);
115 self
116 }
117
118 pub fn set_inbound_limits(self, requests: usize, bytes: usize) -> Self {
137 self.inner.set_inbound_limits(requests, bytes);
138 self
139 }
140
141 pub fn requester(&self) -> Requester {
143 Requester::new(Arc::downgrade(&self.inner))
144 }
145
146 pub fn recv(&mut self) -> Result<(Message, Responder), Error> {
158 self.inner.recv()
159 }
160
161 pub fn closer(&self) -> Closer {
164 Closer::session(Arc::downgrade(&self.inner))
165 }
166
167 pub fn close(&self) {
173 self.inner.close(Error::Closed);
174 }
175}
176
177impl Drop for Session {
178 fn drop(&mut self) {
180 self.close();
181 }
182}
183
184impl fmt::Debug for Session {
185 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
188 let mut session = f.debug_struct("Session");
189 session.field("id", &self.inner.log_id);
190 if let Ok(state) = self.inner.state.try_lock() {
191 session.field("open", &matches!(*state, State::Open { .. }));
192 }
193 session.finish_non_exhaustive()
194 }
195}
196
197pub(super) struct SessionInner {
200 state: Mutex<State>,
202 changed: Condvar,
204 retained_bytes: Arc<AtomicUsize>,
208 side: Side,
210 pub(super) log_id: LogId,
212 stream_closer: Option<transport::Closer>,
215 #[cfg(any(test, feature = "fuzz"))]
217 pub(super) workers: Arc<worker::Tracker>,
218 #[cfg(any(test, feature = "fuzz"))]
220 time: Mutex<Option<Instant>>,
221 #[cfg(any(test, feature = "fuzz"))]
223 drop_hook: Mutex<Option<std::sync::mpsc::Sender<()>>>,
224 #[cfg(any(test, feature = "fuzz"))]
226 disconnect_hook: Mutex<Option<(std::sync::mpsc::Sender<()>, std::sync::mpsc::Receiver<()>)>>,
227}
228
229#[cfg_attr(any(test, feature = "fuzz"), allow(clippy::large_enum_variant))]
233enum State {
234 Open {
237 max_inbound_requests: usize,
239 max_inbound_bytes: usize,
241 autoreply_timeout: Duration,
243
244 incoming: VecDeque<(u64, IncomingEnvelope)>,
246 reserved_ids: HashSet<u64>,
252
253 outgoing: VecDeque<OutgoingMessage>,
255 next_id: Option<u64>,
257 outstanding: HashMap<u64, OperationKey>,
260 operations: HashMap<OperationKey, PendingOperation>,
263
264 #[cfg(any(test, feature = "fuzz"))]
266 wait_hook: Option<std::sync::mpsc::Sender<()>>,
267 },
268 Closed(Error),
270}
271
272impl SessionInner {
273 fn new(
275 side: Side,
276 log_id: LogId,
277 stream_closer: Option<transport::Closer>,
278 #[cfg(any(test, feature = "fuzz"))] workers: Arc<worker::Tracker>,
279 ) -> Self {
280 Self {
281 state: Mutex::new(State::Open {
282 max_inbound_requests: DEFAULT_MAX_INBOUND_REQUESTS,
283 max_inbound_bytes: DEFAULT_MAX_INBOUND_BYTES,
284 autoreply_timeout: DEFAULT_AUTOREPLY_TIMEOUT,
285
286 incoming: VecDeque::new(),
287 reserved_ids: HashSet::new(),
288
289 outgoing: VecDeque::new(),
290 next_id: Some(Parity::from(side).first()),
291 outstanding: HashMap::new(),
292 operations: HashMap::new(),
293
294 #[cfg(any(test, feature = "fuzz"))]
295 wait_hook: None,
296 }),
297 changed: Condvar::new(),
298 retained_bytes: Arc::new(AtomicUsize::new(0)),
299 side,
300 log_id,
301 stream_closer,
302 #[cfg(any(test, feature = "fuzz"))]
303 workers,
304 #[cfg(any(test, feature = "fuzz"))]
305 time: Mutex::new(None),
306 #[cfg(any(test, feature = "fuzz"))]
307 drop_hook: Mutex::new(None),
308 #[cfg(any(test, feature = "fuzz"))]
309 disconnect_hook: Mutex::new(None),
310 }
311 }
312
313 pub(super) fn set_autoreply_timeout(&self, timeout: Duration) {
316 let mut state = self.state.lock().expect("session state not poisoned");
317 if let State::Open {
318 autoreply_timeout, ..
319 } = &mut *state
320 {
321 *autoreply_timeout = timeout;
322 }
323 }
324
325 pub(super) fn set_inbound_limits(&self, requests: usize, bytes: usize) {
328 let removed = {
329 let mut state = self.state.lock().expect("session state not poisoned");
330 let State::Open {
331 max_inbound_requests,
332 max_inbound_bytes,
333 reserved_ids,
334 ..
335 } = &mut *state
336 else {
337 return;
338 };
339 *max_inbound_requests = requests;
340 *max_inbound_bytes = bytes;
341 let error = if reserved_ids.len() > *max_inbound_requests {
342 Some(Error::InboundRequestLimitExceeded(*max_inbound_requests))
343 } else if self.retained_bytes.load(Ordering::Relaxed) > *max_inbound_bytes {
344 Some(Error::InboundByteLimitExceeded(*max_inbound_bytes))
345 } else {
346 None
347 };
348 error.and_then(|error| {
349 let (pending, queued) = state.workload();
350 let removed = state.close(error.clone(), self.now());
351 if removed.is_some() {
352 self.log_closed(&error, pending, queued);
353 }
354 removed
355 })
356 };
357 if removed.is_some() {
358 self.finish_close(removed);
359 }
360 }
361
362 fn retain_incoming(
365 self: &Arc<Self>,
366 bytes: Bytes,
367 header: Header,
368 limit: usize,
369 ) -> Result<IncomingEnvelope, Error> {
370 IncomingEnvelope::new(
371 bytes,
372 header,
373 &self.retained_bytes,
374 limit,
375 self.side,
376 Arc::downgrade(self),
377 )
378 }
379
380 fn recv(self: &Arc<Self>) -> Result<(Message, Responder), Error> {
385 let mut state = self.state.lock().expect("session state not poisoned");
386 let (message, responder) = loop {
387 match &mut *state {
388 State::Closed(error) => return Err(error.clone()),
389 State::Open {
390 incoming,
391 #[cfg(any(test, feature = "fuzz"))]
392 wait_hook,
393 ..
394 } => {
395 if let Some((id, message)) = incoming.pop_front() {
396 break (message, Responder::new(Arc::downgrade(self), id));
397 }
398 #[cfg(any(test, feature = "fuzz"))]
399 if let Some(wait_hook) = wait_hook.take() {
400 let _ = wait_hook.send(());
401 }
402 state = self
403 .changed
404 .wait(state)
405 .expect("session state not poisoned");
406 }
407 }
408 };
409 drop(state);
410 Ok((message.decode()?, responder))
411 }
412
413 pub(super) fn close(&self, error: Error) {
417 let (removed, pending, queued) = {
418 let mut state = self.state.lock().expect("session state not poisoned");
419 let (pending, queued) = state.workload();
420 (state.close(error.clone(), self.now()), pending, queued)
421 };
422 if removed.is_some() {
423 self.log_closed(&error, pending, queued);
424 }
425 self.finish_close(removed);
426 }
427
428 fn log_closed(&self, error: &Error, pending: usize, queued: usize) {
431 match error {
432 Error::Closed => tracing::info!(
433 "wire session {} closed locally (pending {}, queued {})",
434 self.log_id,
435 pending,
436 queued
437 ),
438 _ if error.orderly() => tracing::info!(
439 "wire session {} closed: {} (pending {}, queued {})",
440 self.log_id,
441 error.reason(),
442 pending,
443 queued
444 ),
445 _ => tracing::warn!(
446 "wire session {} failed: {} (pending {}, queued {})",
447 self.log_id,
448 error.reason(),
449 pending,
450 queued
451 ),
452 }
453 }
454
455 fn finish_close(&self, removed: Option<State>) {
457 self.changed.notify_all();
460 drop(removed);
461 if let Some(stream_closer) = &self.stream_closer {
462 stream_closer.close();
463 }
464 }
465
466 pub(super) fn reply_unanswered(self: &Arc<Self>, id: u64) {
469 self.autoreply(
470 id,
471 "unanswered",
472 schema::Error::reserved(
473 schema::ReservedErrors::Unanswered,
474 "request left unanswered",
475 ),
476 );
477 }
478
479 fn reply_unknown(self: &Arc<Self>, id: u64) {
482 self.autoreply(
483 id,
484 "unknown",
485 schema::Error::reserved(schema::ReservedErrors::Unknown, "request not known"),
486 );
487 }
488
489 fn autoreply(self: &Arc<Self>, id: u64, reason: &str, error: schema::Error) {
494 let now = self.now();
495 let timeout = {
496 let state = self.state.lock().expect("session state not poisoned");
497 match &*state {
498 State::Open {
499 autoreply_timeout, ..
500 } => *autoreply_timeout,
501 State::Closed(_) => return,
502 }
503 };
504 tracing::debug!("answering request {} as {}", id, reason);
505
506 let deadline = now.checked_add(timeout).unwrap_or(now);
509 let _ = self.reply(id, Err(error), deadline);
510 }
511
512 pub(super) fn request(
515 self: &Arc<Self>,
516 request: Message,
517 deadline: Instant,
518 ) -> Result<Promise<Message>, Error> {
519 let (sender, promise) = Promise::pair(Arc::downgrade(self), deadline, true);
520 self.enqueue(
521 OutgoingBody::Request(request),
522 PendingOperation {
523 deadline,
524 sender,
525 log_id: None,
526 },
527 )?;
528 Ok(promise)
529 }
530
531 pub(super) fn reply(
534 self: &Arc<Self>,
535 id: u64,
536 result: Result<Message, schema::Error>,
537 deadline: Instant,
538 ) -> Result<Promise<()>, Error> {
539 let (sender, promise) = Promise::pair(Arc::downgrade(self), deadline, false);
540 self.enqueue(
541 OutgoingBody::Reply { id, result },
542 PendingOperation {
543 deadline,
544 sender,
545 log_id: Some(LogId::from(id)),
546 },
547 )?;
548 Ok(promise)
549 }
550
551 fn enqueue(
555 self: &Arc<Self>,
556 body: OutgoingBody,
557 operation: PendingOperation,
558 ) -> Result<(), Error> {
559 {
560 let mut state = self.state.lock().expect("session state not poisoned");
561 let (operations, outgoing, reserved_ids) = match &mut *state {
562 State::Open {
563 operations,
564 outgoing,
565 reserved_ids,
566 ..
567 } => (operations, outgoing, reserved_ids),
568 State::Closed(error) => return Err(error.clone()),
569 };
570 let now = self.now();
571 if now >= operation.deadline {
572 if let OutgoingBody::Reply { id, .. } = body {
573 reserved_ids.remove(&id);
574 }
575 operation.fail(Error::Timeout, now);
576 return Ok(());
577 }
578 let key = OperationKey::new();
579 outgoing.push_back(OutgoingMessage {
580 body,
581 operation: OperationHandle {
582 session: Arc::downgrade(self),
583 key: key.clone(),
584 },
585 #[cfg(any(test, feature = "fuzz"))]
586 deadline: operation.deadline,
587 });
588 operations.insert(key, operation);
589 }
590 self.changed.notify_all();
591 Ok(())
592 }
593
594 pub(super) fn expire(&self) {
598 let mut state = self.state.lock().expect("session state not poisoned");
599 state.expire(self.now());
600 }
601
602 #[cfg(any(test, feature = "fuzz"))]
604 pub(super) fn next_deadline(&self) -> Option<Instant> {
605 let state = self.state.lock().expect("session state not poisoned");
606 match &*state {
607 State::Open { operations, .. } => operations
608 .values()
609 .map(|operation| operation.deadline)
610 .min(),
611 State::Closed(_) => None,
612 }
613 }
614
615 #[cfg(any(test, feature = "fuzz"))]
617 pub(super) fn take_outgoing(&self) -> Option<OutgoingMessage> {
618 let mut state = self.state.lock().expect("session state not poisoned");
619 state.expire(self.now());
620 match &mut *state {
621 State::Open {
622 outgoing,
623 reserved_ids,
624 ..
625 } => {
626 let message = outgoing.pop_front()?;
627 if let OutgoingBody::Reply { id, .. } = &message.body {
628 reserved_ids.remove(id);
629 }
630 Some(message)
631 }
632 State::Closed(_) => None,
633 }
634 }
635
636 pub(super) fn record_write(&self, key: &OperationKey, result: Result<(), Error>) {
639 let mut state = self.state.lock().expect("session state not poisoned");
640 let State::Open { operations, .. } = &mut *state else {
641 return;
642 };
643 let Some(operation) = operations.get(key) else {
644 return;
645 };
646 let now = self.now();
647 if now >= operation.deadline || result.is_err() {
648 let operation = operations.remove(key).expect("operation held under lock");
649 operation.fail(result.err().unwrap_or(Error::Timeout), now);
650 } else if !operation.sender.response {
651 let operation = operations.remove(key).expect("operation held under lock");
652 let _ = operation.sender.send(Ok(PromiseResult::Written));
653 }
654 }
655
656 #[cfg(any(test, feature = "fuzz"))]
659 pub(super) fn record_response(
660 self: &Arc<Self>,
661 key: &OperationKey,
662 result: Result<Message, Error>,
663 ) {
664 let mut state = self.state.lock().expect("session state not poisoned");
665 let State::Open {
666 operations,
667 max_inbound_bytes,
668 ..
669 } = &mut *state
670 else {
671 return;
672 };
673 let Some(operation) = operations.remove(key) else {
674 return;
675 };
676 let result = match result {
677 Ok(message) => Ok(message),
678 Err(Error::Remote(error)) => Err(error),
679 Err(error) => {
680 operation.fail(error, self.now());
681 return;
682 }
683 };
684 let peer = match self.side {
685 Side::Client => Side::Server,
686 Side::Server => Side::Client,
687 };
688 let bytes = peer
689 .encode(0, result)
690 .expect("fixture response belongs to the peer");
691 let bytes = Bytes::from(bytes.into_boxed_slice());
692 let header = self
693 .side
694 .decode_header(bytes.clone())
695 .expect("fixture response has a valid envelope");
696 let result = operation.complete_response(self.now(), || {
697 self.retain_incoming(bytes, header, *max_inbound_bytes)
698 });
699 drop(state);
700 if let Err(error) = result {
701 self.close(error);
702 }
703 }
704
705 fn now(&self) -> Instant {
708 #[cfg(any(test, feature = "fuzz"))]
709 if let Some(now) = *self.time.lock().expect("scenario clock not poisoned") {
710 return now;
711 }
712 Instant::now()
713 }
714
715 pub(super) fn handle_message(self: &Arc<Self>, bytes: Vec<u8>) -> Result<(), Error> {
719 let bytes = Bytes::from(bytes.into_boxed_slice());
720 let header = self.side.decode_header(bytes.clone())?;
721 let mut unknown = None;
722 {
723 let mut state = self.state.lock().expect("session state not poisoned");
724 match MessageKind::from_id(header.id, self.side.into()) {
725 MessageKind::Request => {
726 if header.failed {
727 return Err(self.side.malformed(
728 Some(header),
729 bytes.len(),
730 "envelope",
731 "request contains an error",
732 ));
733 }
734 self.admit_request(&mut state, header, bytes)?;
737 if header.unknown {
738 unknown = Some(header.id);
739 }
740 }
741 MessageKind::Response => {
742 let State::Open {
743 operations,
744 outstanding,
745 max_inbound_bytes,
746 ..
747 } = &mut *state
748 else {
749 let State::Closed(error) = &*state else {
750 unreachable!()
751 };
752 return Err(error.clone());
753 };
754 if let Some(key) = outstanding.remove(&header.id) {
757 if let Some(operation) = operations.remove(&key) {
758 tracing::trace!(
759 "received response {} ({})",
760 header.id,
761 header.payload.unwrap_or("none")
762 );
763 operation.complete_response(self.now(), || {
764 self.retain_incoming(bytes, header, *max_inbound_bytes)
765 })?;
766 } else {
767 tracing::debug!("discarding late response {}", header.id);
768 }
769 } else {
770 tracing::warn!("discarding unmatched response {}", header.id);
771 }
772 }
773 }
774 }
775 if let Some(id) = unknown {
776 self.reply_unknown(id);
777 }
778 self.changed.notify_all();
779 Ok(())
780 }
781
782 fn admit_request(
786 self: &Arc<Self>,
787 state: &mut State,
788 header: Header,
789 bytes: Bytes,
790 ) -> Result<(), Error> {
791 let State::Open {
792 incoming,
793 reserved_ids,
794 max_inbound_requests,
795 max_inbound_bytes,
796 ..
797 } = state
798 else {
799 let State::Closed(error) = state else {
800 unreachable!()
801 };
802 return Err(error.clone());
803 };
804 if reserved_ids.contains(&header.id) {
805 return Err(self.side.malformed(
806 Some(header),
807 bytes.len(),
808 "envelope",
809 "duplicate request ID",
810 ));
811 }
812 let id = header.id;
813 if reserved_ids.len() >= *max_inbound_requests {
814 tracing::warn!(
815 "inbound request limit exceeded (id: {}, used: {}, limit: {})",
816 id,
817 reserved_ids.len(),
818 max_inbound_requests
819 );
820 return Err(Error::InboundRequestLimitExceeded(*max_inbound_requests));
821 }
822 let payload = header.payload.unwrap_or("none");
823 let message = if header.unknown {
824 None
825 } else {
826 Some(self.retain_incoming(bytes, header, *max_inbound_bytes)?)
827 };
828 reserved_ids.insert(id);
829 if let Some(message) = message {
830 incoming.push_back((id, message));
831 }
832 tracing::trace!("received request {} ({})", id, payload);
833 Ok(())
834 }
835
836 pub(super) fn next_outgoing(&self) -> Option<(u64, OutgoingMessage)> {
841 let mut state = self.state.lock().expect("session state not poisoned");
842 loop {
843 state.expire(self.now());
845 let State::Open {
846 outgoing,
847 next_id,
848 outstanding,
849 operations,
850 reserved_ids,
851 ..
852 } = &mut *state
853 else {
854 return None;
855 };
856 if let Some(outgoing) = outgoing.pop_front() {
857 let id = match &outgoing.body {
858 OutgoingBody::Request(_) => {
859 let id = next_id.expect("wire request IDs exhausted");
860 *next_id = id.checked_add(2);
861 outstanding.insert(id, outgoing.operation.key.clone());
864 if let Some(operation) = operations.get_mut(&outgoing.operation.key) {
865 operation.log_id = Some(LogId::from(id));
866 }
867 id
868 }
869 OutgoingBody::Reply { id, .. } => {
870 reserved_ids.remove(id);
874 *id
875 }
876 };
877 return Some((id, outgoing));
878 }
879 state = self
880 .changed
881 .wait(state)
882 .expect("session state not poisoned");
883 }
884 }
885
886 fn run_writer(&self, sender: transport::Sender<impl Write>) {
889 while let Some((id, outgoing)) = self.next_outgoing() {
890 let request = matches!(outgoing.body, OutgoingBody::Request(_));
893 let kind = if request { "request" } else { "reply" };
894 let body = match outgoing.body {
895 OutgoingBody::Request(body) => Ok(body),
896 OutgoingBody::Reply { result, .. } => result,
897 };
898 let payload = match &body {
899 Ok(message) => message.field_name(),
900 Err(_) => "err",
901 };
902 let result = self.side.encode(id, body).and_then(|bytes| {
903 tracing::trace!("sending {} {} ({})", kind, id, payload);
906 sender.send(&bytes).map_err(Error::from)
907 });
908 match &result {
909 Ok(()) => {}
910 Err(Error::Transport(error)) => {
911 self.close(Error::Transport(error.clone()));
914 break;
915 }
916 Err(error) => tracing::debug!("not sending {} {}: {}", kind, id, error),
917 }
918 {
919 let mut state = self.state.lock().expect("session state not poisoned");
922 if let State::Open { outstanding, .. } = &mut *state
923 && request
924 && result.is_err()
925 {
926 outstanding.remove(&id);
927 }
928 }
929 outgoing.operation.record_write(result);
932 }
933 #[cfg(any(test, feature = "fuzz"))]
934 if let Some((entered, released)) = self.disconnect_hook.lock().unwrap().take() {
935 let _ = entered.send(());
936 let _ = released.recv();
937 }
938 if let Err(error) = sender.disconnect() {
941 tracing::debug!(
942 "failed to signal dropped session {}: {}",
943 sender.log_id(),
944 error
945 );
946 }
947 }
948
949 fn run_deadlines(&self) {
952 let mut state = self.state.lock().expect("session state not poisoned");
953 loop {
954 state.expire(self.now());
955 let State::Open { operations, .. } = &*state else {
956 return;
957 };
958 state = match operations
961 .values()
962 .map(|operation| operation.deadline)
963 .min()
964 {
965 Some(deadline) => {
966 self.changed
967 .wait_timeout(state, deadline.saturating_duration_since(self.now()))
968 .expect("session state not poisoned")
969 .0
970 }
971 None => self
972 .changed
973 .wait(state)
974 .expect("session state not poisoned"),
975 };
976 }
977 }
978}
979
980impl Session {
981 pub(super) fn start<W: Write + Send + 'static>(
984 side: Side,
985 sender: transport::Sender<W>,
986 stream_closer: Option<transport::Closer>,
987 #[cfg(any(test, feature = "fuzz"))] workers: Arc<worker::Tracker>,
988 ) -> Self {
989 let session = Self {
990 inner: Arc::new(SessionInner::new(
991 side,
992 sender.log_id(),
993 stream_closer,
994 #[cfg(any(test, feature = "fuzz"))]
995 workers.clone(),
996 )),
997 };
998 let inner = session.inner.clone();
999 worker::spawn(
1000 "wire-writer",
1001 #[cfg(any(test, feature = "fuzz"))]
1002 &workers,
1003 move || inner.run_writer(sender),
1004 );
1005 let inner = session.inner.clone();
1006 worker::spawn(
1007 "wire-deadlines",
1008 #[cfg(any(test, feature = "fuzz"))]
1009 &workers,
1010 move || inner.run_deadlines(),
1011 );
1012 session
1013 }
1014}
1015
1016impl State {
1017 fn workload(&self) -> (usize, usize) {
1019 match self {
1020 Self::Open {
1021 operations,
1022 outgoing,
1023 incoming,
1024 ..
1025 } => (operations.len(), outgoing.len() + incoming.len()),
1026 Self::Closed(_) => (0, 0),
1027 }
1028 }
1029
1030 fn close(&mut self, error: Error, now: Instant) -> Option<Self> {
1033 if let Self::Closed(_) = self {
1034 return None;
1035 }
1036 let mut removed = std::mem::replace(self, Self::Closed(error.clone()));
1037 if let Self::Open { operations, .. } = &mut removed {
1038 for (_, operation) in operations.drain() {
1039 operation.fail(error.clone(), now);
1040 }
1041 }
1042 Some(removed)
1043 }
1044
1045 fn expire(&mut self, now: Instant) {
1048 if let Self::Open {
1049 operations,
1050 outgoing,
1051 reserved_ids,
1052 ..
1053 } = self
1054 {
1055 let expired: Vec<_> = operations
1056 .iter()
1057 .filter(|(_, operation)| now >= operation.deadline)
1058 .map(|(key, _)| key.clone())
1059 .collect();
1060 for key in expired {
1061 operations
1062 .remove(&key)
1063 .expect("expired operation held under lock")
1064 .fail(Error::Timeout, now);
1065 }
1066 outgoing.retain(|outgoing| {
1069 let retained = operations.contains_key(&outgoing.operation.key);
1070 if !retained && let OutgoingBody::Reply { id, .. } = outgoing.body {
1071 reserved_ids.remove(&id);
1072 }
1073 retained
1074 });
1075 }
1076 }
1077}
1078
1079#[cfg(any(test, feature = "fuzz"))]
1081impl Session {
1082 pub(super) fn fixture() -> Self {
1084 Self::fixture_for(Side::Server)
1085 }
1086
1087 pub(super) fn fixture_for(side: Side) -> Self {
1089 Self {
1090 inner: Arc::new(SessionInner::new(
1091 side,
1092 LogId::default(),
1093 None,
1094 Arc::new(worker::Tracker::default()),
1095 )),
1096 }
1097 }
1098}
1099
1100#[cfg(any(test, feature = "fuzz"))]
1101impl SessionInner {
1102 pub(super) fn pause_disconnect(
1105 &self,
1106 ) -> (std::sync::mpsc::Receiver<()>, std::sync::mpsc::Sender<()>) {
1107 let (entered, observed) = std::sync::mpsc::channel();
1108 let (release, released) = std::sync::mpsc::channel();
1109 *self.disconnect_hook.lock().unwrap() = Some((entered, released));
1110 (observed, release)
1111 }
1112
1113 pub(super) fn watch_drop(&self) -> std::sync::mpsc::Receiver<()> {
1115 let (sender, receiver) = std::sync::mpsc::channel();
1116 *self.drop_hook.lock().unwrap() = Some(sender);
1117 receiver
1118 }
1119
1120 pub(super) fn use_last_request_id(&self) {
1122 let mut state = self.state.lock().unwrap();
1123 let State::Open {
1124 next_id,
1125 outstanding,
1126 ..
1127 } = &mut *state
1128 else {
1129 panic!("open session required")
1130 };
1131 assert!(outstanding.is_empty());
1132 *next_id = Some(if self.side == Side::Client {
1133 u64::MAX
1134 } else {
1135 u64::MAX - 1
1136 });
1137 }
1138
1139 pub(super) fn wait_response(&self, id: u64) {
1142 let state = self.state.lock().unwrap();
1143 let (state, _) = self.changed.wait_timeout_while(state, Duration::from_secs(3), |state| {
1144 matches!(state, State::Open { outstanding, .. } if outstanding.contains_key(&id))
1145 }).unwrap();
1146 let State::Open { outstanding, .. } = &*state else {
1147 panic!("session closed before response fence");
1148 };
1149 assert!(
1150 !outstanding.contains_key(&id),
1151 "reader did not process response"
1152 );
1153 }
1154
1155 pub(super) fn outstanding_ids(&self) -> Vec<u64> {
1157 let state = self.state.lock().unwrap();
1158 let State::Open { outstanding, .. } = &*state else {
1159 panic!("open session required")
1160 };
1161 let mut ids: Vec<_> = outstanding.keys().copied().collect();
1162 ids.sort_unstable();
1163 ids
1164 }
1165
1166 pub(super) fn inject_request(self: &Arc<Self>, id: u64, message: Message) -> Result<(), Error> {
1169 let peer = match self.side {
1170 Side::Client => Side::Server,
1171 Side::Server => Side::Client,
1172 };
1173 let bytes = peer
1174 .encode(id, Ok(message))
1175 .expect("fixture request belongs to peer");
1176 let bytes = Bytes::from(bytes.into_boxed_slice());
1177 let header = self.side.decode_header(bytes.clone())?;
1178 let result = {
1179 let mut state = self.state.lock().expect("session state not poisoned");
1180 self.admit_request(&mut state, header, bytes)
1181 };
1182 self.changed.notify_all();
1183 if let Err(error) = &result {
1184 self.close(error.clone());
1185 }
1186 result
1187 }
1188
1189 pub(super) fn inbound_usage(&self) -> (usize, usize) {
1191 let state = self.state.lock().expect("session state not poisoned");
1192 let requests = match &*state {
1193 State::Open { reserved_ids, .. } => reserved_ids.len(),
1194 State::Closed(_) => 0,
1195 };
1196 (requests, self.retained_bytes.load(Ordering::Relaxed))
1197 }
1198
1199 pub(super) fn set_time(&self, now: Instant) {
1202 let _state = self.state.lock().expect("session state not poisoned");
1203 let mut time = self.time.lock().expect("scenario clock not poisoned");
1204 assert!(
1205 time.is_none_or(|previous| now >= previous),
1206 "clock cannot go backwards"
1207 );
1208 *time = Some(now);
1209 }
1210
1211 pub(super) fn use_realtime(&self) {
1213 let _state = self.state.lock().expect("session state not poisoned");
1214 *self.time.lock().expect("scenario clock not poisoned") = None;
1215 }
1216
1217 pub(super) fn watch_recv_wait(&self) -> std::sync::mpsc::Receiver<()> {
1224 let (sender, receiver) = std::sync::mpsc::channel();
1225 let mut state = self.state.lock().expect("session state not poisoned");
1226 let State::Open {
1227 incoming,
1228 wait_hook,
1229 ..
1230 } = &mut *state
1231 else {
1232 panic!("only watch an open session receive");
1233 };
1234 assert!(incoming.is_empty());
1235 *wait_hook = Some(sender);
1236 receiver
1237 }
1238}
1239
1240#[cfg(any(test, feature = "fuzz"))]
1241impl Drop for SessionInner {
1242 fn drop(&mut self) {
1244 if let Some(sender) = self.drop_hook.get_mut().unwrap().take() {
1245 let _ = sender.send(());
1246 }
1247 }
1248}
1249
1250#[cfg(test)]
1252#[cfg_attr(coverage_nightly, coverage(off))]
1253mod tests {
1254 use crate::protocol::{self, Error, Session};
1255 use crate::transport::{Read, Stream, Verifier, Write};
1256 use std::fmt::Debug;
1257
1258 #[allow(dead_code)]
1260 fn connect<R, W, V>(stream: Stream<R, W>, verifier: &V) -> Result<(Session, V::Info), Error>
1261 where
1262 R: Read + Send + 'static,
1263 W: Write + Send + 'static,
1264 V: Verifier,
1265 {
1266 protocol::connect(stream, verifier)
1267 }
1268
1269 #[test]
1272 fn test_thread_capabilities() {
1273 fn movable<T: Debug + Send + 'static>() {}
1275 movable::<Session>();
1276 }
1277}