1use std::collections::HashMap;
14use std::sync::Arc;
15use std::sync::atomic::{AtomicBool, Ordering};
16
17use bytes::Bytes;
18use prost::Message;
19use tokio::sync::{Mutex, OwnedSemaphorePermit, Semaphore, mpsc, oneshot};
20
21use crate::bus::packet::{Packet, PacketFlags, PacketType};
22use crate::bus::{Bus, BusReader, BusWriter};
23use crate::error::{Error, Result};
24use crate::guid::Guid;
25use crate::proto;
26use crate::rpc::{self, ResponseMessage};
27
28const OUTBOUND_QUEUE: usize = 64;
30
31const MAX_IN_FLIGHT: usize = 256;
38
39const CANCEL_QUEUE: usize = MAX_IN_FLIGHT;
46
47#[derive(Debug, Default)]
55struct Waiters {
56 closed: bool,
57 by_request: HashMap<Guid, oneshot::Sender<ResponseMessage>>,
58}
59
60impl Waiters {
61 fn close(&mut self) {
65 self.closed = true;
66 self.by_request.clear();
67 }
68}
69
70type Pending = Arc<Mutex<Waiters>>;
71type InFlight = Arc<Semaphore>;
72
73#[derive(Debug)]
80struct Cancellation {
81 packet: Packet,
82 _permit: OwnedSemaphorePermit,
83}
84
85enum Outgoing {
86 Request(Packet),
87 Cancellation(Cancellation),
88}
89
90#[derive(Debug)]
98pub struct Connection {
99 outbound: mpsc::Sender<Packet>,
100 cancels: mpsc::Sender<Cancellation>,
101 pending: Pending,
102 in_flight: InFlight,
103 address: String,
104 token: Option<String>,
105 closed: Arc<AtomicBool>,
106 reader_task: tokio::task::JoinHandle<()>,
107}
108
109impl Drop for Connection {
110 fn drop(&mut self) {
111 self.reader_task.abort();
118 }
119}
120
121impl Connection {
122 pub async fn connect(address: &str, token: Option<String>) -> Result<Self> {
124 let bus = Bus::connect(address).await?;
125 Ok(Self::from_bus(bus, address.to_owned(), token))
126 }
127
128 fn from_bus(bus: Bus, address: String, token: Option<String>) -> Self {
129 let Bus { reader, writer, .. } = bus;
130 let pending: Pending = Arc::default();
131 let in_flight = Arc::new(Semaphore::new(MAX_IN_FLIGHT));
132 let closed = Arc::new(AtomicBool::new(false));
133 let (outbound, outbound_receiver) = mpsc::channel(OUTBOUND_QUEUE);
134 let (cancels, cancel_receiver) = mpsc::channel(CANCEL_QUEUE);
135
136 tokio::spawn(write_loop(
137 writer,
138 outbound_receiver,
139 cancel_receiver,
140 Arc::clone(&pending),
141 Arc::clone(&in_flight),
142 Arc::clone(&closed),
143 ));
144 let reader_task = tokio::spawn(read_loop(
145 reader,
146 Arc::clone(&pending),
147 Arc::clone(&in_flight),
148 Arc::clone(&closed),
149 ));
150
151 Self {
152 outbound,
153 cancels,
154 pending,
155 in_flight,
156 address,
157 token,
158 closed,
159 reader_task,
160 }
161 }
162
163 pub fn address(&self) -> &str {
165 &self.address
166 }
167
168 pub fn is_closed(&self) -> bool {
170 self.closed.load(Ordering::Relaxed)
171 }
172
173 pub async fn invoke<Response: Message + Default>(
180 &self,
181 method: &str,
182 body: &impl Message,
183 attachments: Vec<Bytes>,
184 timeout: Option<std::time::Duration>,
185 response_name: &'static str,
186 ) -> Result<(Response, Vec<Bytes>)> {
187 let response = self
188 .invoke_raw(rpc::API_SERVICE, method, body, attachments, timeout, None)
189 .await?;
190 let decoded = response.decode_body::<Response>(response_name)?;
191 Ok((decoded, response.attachments))
192 }
193
194 pub async fn invoke_raw(
196 &self,
197 service: &str,
198 method: &str,
199 body: &impl Message,
200 attachments: Vec<Bytes>,
201 timeout: Option<std::time::Duration>,
202 mutation_id: Option<Guid>,
203 ) -> Result<ResponseMessage> {
204 let mut builder = rpc::RequestHeaderBuilder::new(service, method);
205 builder.timeout = timeout;
206 builder.mutation_id = mutation_id;
207 let request_id = builder.request_id;
208 let header = builder.build();
209
210 let deadline = timeout.map(|limit| tokio::time::Instant::now() + limit);
216 let timed_out = || Error::Timeout {
217 service: service.to_owned(),
218 method: method.to_owned(),
219 timeout: timeout.unwrap_or_default(),
220 };
221
222 let permit = match deadline {
226 Some(deadline) => {
227 match tokio::time::timeout_at(deadline, Arc::clone(&self.in_flight).acquire_owned())
228 .await
229 {
230 Ok(Ok(permit)) => permit,
231 Ok(Err(_)) => return Err(Error::ConnectionClosed { request_id }),
232 Err(_) => return Err(timed_out()),
233 }
234 }
235 None => Arc::clone(&self.in_flight)
236 .acquire_owned()
237 .await
238 .map_err(|_| Error::ConnectionClosed { request_id })?,
239 };
240
241 let (sender, receiver) = oneshot::channel();
242 {
243 let mut waiters = self.pending.lock().await;
244 if waiters.closed {
248 return Err(Error::ConnectionClosed { request_id });
249 }
250 waiters.by_request.insert(request_id, sender);
251 }
252
253 let mut guard = PendingGuard {
258 pending: Arc::clone(&self.pending),
259 cancels: self.cancels.clone(),
260 request_id,
261 service: service.to_owned(),
262 method: method.to_owned(),
263 completed: false,
264 sent: false,
265 permit: Some(permit),
266 };
267
268 let parts = rpc::encode_request(&header, self.token.as_deref(), body, attachments);
269 let packet = Packet::message(Guid::random(), parts, PacketFlags::NONE);
270 let queued = match deadline {
271 Some(deadline) => {
272 match tokio::time::timeout_at(deadline, self.outbound.send(packet)).await {
273 Ok(queued) => queued,
274 Err(_) => return Err(timed_out()),
277 }
278 }
279 None => self.outbound.send(packet).await,
280 };
281 if queued.is_err() {
282 return Err(Error::ConnectionClosed { request_id });
286 }
287 guard.sent = true;
288
289 let response = match deadline {
290 Some(deadline) => match tokio::time::timeout_at(deadline, receiver).await {
291 Ok(received) => received,
292 Err(_) => return Err(timed_out()),
295 },
296 None => receiver.await,
297 };
298
299 let response = match response {
300 Ok(response) => response,
301 Err(_) => {
304 guard.complete();
305 return Err(Error::ConnectionClosed { request_id });
306 }
307 };
308 guard.complete();
310
311 if let Some(error) = response.error() {
312 return Err(Error::response(service, method, error));
313 }
314 Ok(response)
315 }
316}
317
318struct PendingGuard {
330 pending: Pending,
331 cancels: mpsc::Sender<Cancellation>,
332 request_id: Guid,
333 service: String,
334 method: String,
335 completed: bool,
337 sent: bool,
342 permit: Option<OwnedSemaphorePermit>,
344}
345
346impl PendingGuard {
347 fn complete(&mut self) {
348 self.completed = true;
349 }
350}
351
352impl Drop for PendingGuard {
353 fn drop(&mut self) {
354 if self.completed {
355 return;
356 }
357
358 let pending = Arc::clone(&self.pending);
359 let request_id = self.request_id;
360 match tokio::runtime::Handle::try_current() {
368 Ok(handle) => {
369 handle.spawn(async move {
370 pending.lock().await.by_request.remove(&request_id);
371 });
372 }
373 Err(_) => {
374 if let Ok(mut waiters) = pending.try_lock() {
375 waiters.by_request.remove(&request_id);
376 }
377 }
383 }
384
385 if !self.sent {
386 return;
387 }
388
389 let parts = rpc::encode_cancelation(request_id, &self.service, &self.method);
394 let cancellation = Cancellation {
395 packet: Packet::message(Guid::random(), parts, PacketFlags::NONE),
396 _permit: self
397 .permit
398 .take()
399 .expect("every unfinished call holds an in-flight permit"),
400 };
401 match self.cancels.try_send(cancellation) {
402 Ok(()) | Err(mpsc::error::TrySendError::Closed(_)) => {}
403 Err(mpsc::error::TrySendError::Full(_)) => {
408 debug_assert!(false, "cancellation queue exceeded in-flight limit");
409 }
410 }
411 }
412}
413
414async fn write_loop(
415 mut writer: BusWriter,
416 mut outbound: mpsc::Receiver<Packet>,
417 mut cancels: mpsc::Receiver<Cancellation>,
418 pending: Pending,
419 in_flight: InFlight,
420 closed: Arc<AtomicBool>,
421) {
422 loop {
423 let outgoing = tokio::select! {
428 biased;
429 Some(cancellation) = cancels.recv() => Outgoing::Cancellation(cancellation),
430 Some(packet) = outbound.recv() => Outgoing::Request(packet),
431 else => break,
432 };
433 let packet = match &outgoing {
434 Outgoing::Request(packet) => packet,
435 Outgoing::Cancellation(cancellation) => &cancellation.packet,
436 };
437 if writer.send(packet).await.is_err() {
438 break;
439 }
440 }
441 closed.store(true, Ordering::Relaxed);
442 in_flight.close();
443 pending.lock().await.close();
444 let _ = writer.shutdown().await;
445}
446
447async fn read_loop(
448 mut reader: BusReader,
449 pending: Pending,
450 in_flight: InFlight,
451 closed: Arc<AtomicBool>,
452) {
453 loop {
454 let packet = match reader.receive().await {
455 Ok(packet) => packet,
456 Err(_) => break,
457 };
458
459 if packet.packet_type != PacketType::Message {
462 continue;
463 }
464
465 let Ok(response) = rpc::decode_response(packet.parts) else {
466 continue;
469 };
470 let Some(request_id) = response.request_id() else {
471 continue;
472 };
473 if let Some(sender) = pending.lock().await.by_request.remove(&request_id) {
474 let _ = sender.send(response);
475 }
476 }
477
478 closed.store(true, Ordering::Relaxed);
479 in_flight.close();
480 pending.lock().await.close();
484}
485
486pub async fn discover_proxies(
492 connection: &Connection,
493 role: Option<&str>,
494 timeout: Option<std::time::Duration>,
495) -> Result<Vec<String>> {
496 let request = proto::api::TReqDiscoverProxies {
497 role: role.map(str::to_owned),
498 ..Default::default()
499 };
500 let response = connection
501 .invoke_raw(
502 rpc::DISCOVERY_SERVICE,
503 "DiscoverProxies",
504 &request,
505 Vec::new(),
506 timeout,
507 None,
508 )
509 .await?;
510 let decoded = response.decode_body::<proto::api::TRspDiscoverProxies>("TRspDiscoverProxies")?;
511 Ok(decoded.addresses)
512}
513
514#[cfg(test)]
515mod tests {
516 use super::*;
517 use crate::bus::packet;
518 use bytes::BytesMut;
519 use tokio::io::{AsyncReadExt, AsyncWriteExt};
520 use tokio::net::TcpListener;
521
522 struct StubProxy {
531 address: String,
532 seen: mpsc::UnboundedReceiver<Packet>,
533 task: tokio::task::JoinHandle<()>,
534 inject: mpsc::UnboundedSender<Packet>,
535 }
536
537 impl StubProxy {
538 async fn inject(&self, packet: Packet) {
540 let _ = self.inject.send(packet);
541 tokio::time::sleep(std::time::Duration::from_millis(20)).await;
543 }
544 }
545
546 impl Drop for StubProxy {
547 fn drop(&mut self) {
548 self.task.abort();
549 }
550 }
551
552 async fn stub_proxy(
553 answer: impl Fn(&proto::rpc::TRequestHeader) -> Option<Vec<Option<Bytes>>> + Send + 'static,
554 ) -> StubProxy {
555 stub_proxy_with_batching(answer, 1).await
556 }
557
558 async fn stub_proxy_with_batching(
567 answer: impl Fn(&proto::rpc::TRequestHeader) -> Option<Vec<Option<Bytes>>> + Send + 'static,
568 batch: usize,
569 ) -> StubProxy {
570 let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
571 let address = listener.local_addr().unwrap().to_string();
572 let (seen_sender, seen) = mpsc::unbounded_channel();
573 let (inject, mut injected) = mpsc::unbounded_channel::<Packet>();
574
575 let task = tokio::spawn(async move {
576 let (stream, _) = listener.accept().await.unwrap();
577 let (mut read_half, mut write_half) = stream.into_split();
578 let mut buffer = BytesMut::new();
579 let mut handshaken = false;
580 let mut pending_replies: Vec<Vec<Option<Bytes>>> = Vec::new();
581
582 loop {
583 while let Ok(packet) = injected.try_recv() {
585 let mut out = BytesMut::new();
586 packet::encode(&packet, &mut out).unwrap();
587 if write_half.write_all(&out).await.is_err() {
588 return;
589 }
590 }
591
592 let decoded = packet::decode(&mut buffer, crate::bus::DEFAULT_MAX_MESSAGE_SIZE);
593 match decoded {
594 Ok(Some(request)) => {
595 if !handshaken {
596 handshaken = true;
597 let handshake = proto::bus::THandshake {
598 connection_id: Guid::random().to_proto(),
599 encryption_mode: Some(0),
600 ..Default::default()
601 };
602 let mut part = Vec::new();
603 part.extend_from_slice(&crate::bus::HANDSHAKE_SIGNATURE.to_le_bytes());
604 handshake.encode(&mut part).unwrap();
605 let reply = Packet::message(
606 request.id,
607 vec![Some(Bytes::from(part))],
608 PacketFlags::NONE,
609 );
610 let mut out = BytesMut::new();
611 packet::encode(&reply, &mut out).unwrap();
612 if write_half.write_all(&out).await.is_err() {
613 return;
614 }
615 continue;
616 }
617
618 let Some(Some(header_part)) = request.parts.first().cloned() else {
623 continue;
624 };
625 let _ = seen_sender.send(request.clone());
626 if header_part.len() < 4 {
627 continue;
628 }
629 let Ok(header) = proto::rpc::TRequestHeader::decode(&header_part[4..])
630 else {
631 continue;
632 };
633
634 if let Some(parts) = answer(&header) {
635 pending_replies.push(parts);
636 }
637 if pending_replies.len() >= batch {
638 for parts in pending_replies.drain(..).rev() {
641 let reply =
642 Packet::message(Guid::random(), parts, PacketFlags::NONE);
643 let mut out = BytesMut::new();
644 packet::encode(&reply, &mut out).unwrap();
645 if write_half.write_all(&out).await.is_err() {
646 return;
647 }
648 }
649 }
650 continue;
651 }
652 Ok(None) => {}
653 Err(_) => return,
654 }
655 if read_half.read_buf(&mut buffer).await.unwrap_or(0) == 0 {
656 return;
657 }
658 }
659 });
660
661 StubProxy {
662 address,
663 seen,
664 task,
665 inject,
666 }
667 }
668
669 async fn next_packet(stub: &mut StubProxy) -> Option<Packet> {
675 tokio::time::timeout(std::time::Duration::from_secs(5), stub.seen.recv())
676 .await
677 .ok()
678 .flatten()
679 }
680
681 fn success_reply(request_id: Guid, body: &impl Message) -> Vec<Option<Bytes>> {
682 let header = proto::rpc::TResponseHeader {
683 request_id: Some(request_id.to_proto()),
684 ..Default::default()
685 };
686 let mut header_part = Vec::new();
687 header_part.extend_from_slice(&(rpc::MessageType::Response as u32).to_le_bytes());
688 header.encode(&mut header_part).unwrap();
689 vec![
690 Some(Bytes::from(header_part)),
691 Some(Bytes::from(body.encode_to_vec())),
692 ]
693 }
694
695 fn error_reply(request_id: Guid, code: i32, message: &str) -> Vec<Option<Bytes>> {
696 let header = proto::rpc::TResponseHeader {
697 request_id: Some(request_id.to_proto()),
698 error: Some(proto::misc::TError {
699 code,
700 message: Some(message.to_owned()),
701 attributes: None,
702 inner_errors: vec![],
703 }),
704 ..Default::default()
705 };
706 let mut header_part = Vec::new();
707 header_part.extend_from_slice(&(rpc::MessageType::Response as u32).to_le_bytes());
708 header.encode(&mut header_part).unwrap();
709 vec![Some(Bytes::from(header_part))]
710 }
711
712 #[tokio::test]
713 async fn a_call_gets_its_own_response() {
714 let mut stub = stub_proxy(|header| {
715 let request_id = Guid::from_proto(header.request_id.as_ref().unwrap());
716 Some(success_reply(
717 request_id,
718 &proto::api::TRspPingTransaction::default(),
719 ))
720 })
721 .await;
722
723 let connection = Connection::connect(&stub.address, None).await.unwrap();
724 let request = proto::api::TReqPingTransaction {
725 transaction_id: Guid::random().to_proto(),
726 ..Default::default()
727 };
728 let (_response, attachments) = tokio::time::timeout(
732 std::time::Duration::from_secs(10),
733 connection.invoke::<proto::api::TRspPingTransaction>(
734 "PingTransaction",
735 &request,
736 Vec::new(),
737 None,
738 "TRspPingTransaction",
739 ),
740 )
741 .await
742 .expect("the stub answers immediately")
743 .unwrap();
744 assert!(attachments.is_empty(), "the stub sent no attachments");
745
746 let sent = next_packet(&mut stub).await.expect("the request");
750 let header_part = sent.parts[0].as_ref().unwrap();
751 let header = proto::rpc::TRequestHeader::decode(&header_part[4..]).unwrap();
752 assert_eq!(header.method, "PingTransaction");
753 assert_eq!(header.service, rpc::API_SERVICE);
754 let body = proto::api::TReqPingTransaction::decode(sent.parts[1].as_ref().unwrap().clone())
755 .unwrap();
756 assert_eq!(body.transaction_id, request.transaction_id);
757 }
758
759 #[tokio::test]
769 async fn concurrent_requests_are_routed_by_request_id() {
770 let stub = stub_proxy_with_batching(
771 |header| {
772 let request_id = Guid::from_proto(header.request_id.as_ref().unwrap());
773 Some(success_reply(
776 request_id,
777 &proto::api::TRspGetNode {
778 value: header.method.clone().into_bytes(),
779 },
780 ))
781 },
782 4,
783 )
784 .await;
785
786 let connection = Arc::new(Connection::connect(&stub.address, None).await.unwrap());
787 let methods = ["GetNode", "ListNode", "ExistsNode", "SetNode"];
788 let mut handles = Vec::new();
789 for method in methods {
790 let connection = Arc::clone(&connection);
791 handles.push(tokio::spawn(async move {
792 connection
793 .invoke::<proto::api::TRspGetNode>(
794 method,
795 &proto::api::TReqGetNode::default(),
796 Vec::new(),
797 None,
798 "TRspGetNode",
799 )
800 .await
801 .map(|(response, _)| String::from_utf8(response.value).unwrap())
802 }));
803 }
804
805 for (method, handle) in methods.iter().zip(handles) {
806 let answer = tokio::time::timeout(std::time::Duration::from_secs(10), handle)
807 .await
808 .expect("a call is stuck: the stub answers only once all four have arrived")
809 .unwrap()
810 .unwrap();
811 assert_eq!(
812 &answer, method,
813 "the caller for {method} was handed another call's answer"
814 );
815 }
816 }
817
818 #[tokio::test]
819 async fn a_server_error_becomes_a_rust_error_with_its_code() {
820 let stub = stub_proxy(|header| {
821 let request_id = Guid::from_proto(header.request_id.as_ref().unwrap());
822 Some(error_reply(
823 request_id,
824 crate::error::codes::NO_SUCH_TRANSACTION,
825 "no such transaction",
826 ))
827 })
828 .await;
829
830 let connection = Connection::connect(&stub.address, None).await.unwrap();
831 let error = tokio::time::timeout(
832 std::time::Duration::from_secs(10),
833 connection.invoke::<proto::api::TRspPingTransaction>(
834 "PingTransaction",
835 &proto::api::TReqPingTransaction {
836 transaction_id: Guid::random().to_proto(),
837 ..Default::default()
838 },
839 Vec::new(),
840 None,
841 "TRspPingTransaction",
842 ),
843 )
844 .await
845 .expect("the stub answers immediately")
846 .unwrap_err();
847
848 assert!(error.has_code(crate::error::codes::NO_SUCH_TRANSACTION));
849 assert!(
850 error
851 .to_string()
852 .contains("ApiService.PingTransaction failed")
853 );
854 }
855
856 #[tokio::test]
857 async fn a_timeout_reports_the_method_and_cancels_the_request() {
858 let mut stub = stub_proxy(|_| None).await;
860 let connection = Connection::connect(&stub.address, None).await.unwrap();
861
862 let error = tokio::time::timeout(
867 std::time::Duration::from_secs(10),
868 connection.invoke::<proto::api::TRspPingTransaction>(
869 "PingTransaction",
870 &proto::api::TReqPingTransaction {
871 transaction_id: Guid::random().to_proto(),
872 ..Default::default()
873 },
874 Vec::new(),
875 Some(std::time::Duration::from_millis(50)),
876 "TRspPingTransaction",
877 ),
878 )
879 .await
880 .expect("the local deadline did not fire: the call outlived it twentyfold")
881 .unwrap_err();
882 assert!(matches!(error, Error::Timeout { .. }), "got {error}");
883
884 let request = next_packet(&mut stub).await.expect("the request");
886 let header_part = request.parts[0].as_ref().unwrap();
887 assert_eq!(&header_part[0..4], b"rpci");
888 let header = proto::rpc::TRequestHeader::decode(&header_part[4..]).unwrap();
889 let request_id = Guid::from_proto(header.request_id.as_ref().unwrap());
890 assert_eq!(header.timeout, Some(50_000));
893
894 let cancelation = next_packet(&mut stub)
895 .await
896 .expect("a cancellation must follow the timeout");
897 let part = cancelation.parts[0].as_ref().unwrap();
898 assert_eq!(&part[0..4], b"rpcc", "cancellation is an rpcc message");
899 let cancel_header = proto::rpc::TRequestCancelationHeader::decode(&part[4..]).unwrap();
900 assert_eq!(Guid::from_proto(&cancel_header.request_id), request_id);
901 }
902
903 #[tokio::test]
910 async fn junk_from_the_peer_does_not_kill_the_connection() {
911 let stub = stub_proxy(|header| {
912 let request_id = Guid::from_proto(header.request_id.as_ref().unwrap());
913 Some(success_reply(
914 request_id,
915 &proto::api::TRspPingTransaction::default(),
916 ))
917 })
918 .await;
919
920 let connection = Connection::connect(&stub.address, None).await.unwrap();
921 let request = proto::api::TReqPingTransaction {
922 transaction_id: Guid::random().to_proto(),
923 ..Default::default()
924 };
925
926 let orphan = {
929 let header = proto::rpc::TResponseHeader {
930 request_id: Some(Guid::random().to_proto()),
931 ..Default::default()
932 };
933 let mut bytes = Vec::new();
934 bytes.extend_from_slice(&(rpc::MessageType::Response as u32).to_le_bytes());
935 header.encode(&mut bytes).unwrap();
936 Packet::message(
937 Guid::random(),
938 vec![Some(Bytes::from(bytes))],
939 PacketFlags::NONE,
940 )
941 };
942 let unparseable = Packet::message(
943 Guid::random(),
944 vec![Some(Bytes::from_static(b"not an rpc message at all"))],
945 PacketFlags::NONE,
946 );
947 let ack = Packet {
948 packet_type: PacketType::Ack,
949 flags: PacketFlags::NONE,
950 id: Guid::random(),
951 parts: Vec::new(),
952 };
953
954 for packet in [orphan, unparseable, ack] {
955 stub.inject(packet).await;
956 }
957 tokio::time::sleep(std::time::Duration::from_millis(100)).await;
958
959 assert!(!connection.is_closed(), "junk closed the connection");
961 tokio::time::timeout(
962 std::time::Duration::from_secs(10),
963 connection.invoke::<proto::api::TRspPingTransaction>(
964 "PingTransaction",
965 &request,
966 Vec::new(),
967 Some(std::time::Duration::from_secs(5)),
968 "TRspPingTransaction",
969 ),
970 )
971 .await
972 .expect("the connection stopped answering after the junk")
973 .expect("a call after the junk must still work");
974 }
975
976 #[tokio::test]
977 async fn a_dropped_connection_fails_the_calls_in_flight() {
978 let stub = stub_proxy(|_| None).await;
979 let connection = Connection::connect(&stub.address, None).await.unwrap();
980
981 let request = proto::api::TReqPingTransaction {
982 transaction_id: Guid::random().to_proto(),
983 ..Default::default()
984 };
985 let call = connection.invoke::<proto::api::TRspPingTransaction>(
986 "PingTransaction",
987 &request,
988 Vec::new(),
989 None,
990 "TRspPingTransaction",
991 );
992
993 drop(stub);
996 let error = tokio::time::timeout(std::time::Duration::from_secs(10), call)
997 .await
998 .expect("dropping the connection must fail the call, not park it")
999 .unwrap_err();
1000 assert!(
1001 matches!(error, Error::ConnectionClosed { .. }),
1002 "got {error}"
1003 );
1004 }
1005
1006 #[tokio::test]
1011 async fn dropping_a_call_cancels_it_on_the_wire() {
1012 let mut stub = stub_proxy(|_| None).await;
1013 let connection = Connection::connect(&stub.address, None).await.unwrap();
1014
1015 let request = proto::api::TReqSelectRows {
1016 query: "* from [//tmp/t]".to_owned(),
1017 ..Default::default()
1018 };
1019 {
1020 let call = connection.invoke::<proto::api::TRspSelectRows>(
1021 "SelectRows",
1022 &request,
1023 Vec::new(),
1024 None,
1026 "TRspSelectRows",
1027 );
1028 let _ = tokio::time::timeout(std::time::Duration::from_millis(50), call).await;
1029 }
1030
1031 let sent = next_packet(&mut stub).await.expect("the request");
1032 let header_part = sent.parts[0].as_ref().unwrap();
1033 assert_eq!(&header_part[0..4], b"rpci");
1034 let header = proto::rpc::TRequestHeader::decode(&header_part[4..]).unwrap();
1035 let request_id = Guid::from_proto(header.request_id.as_ref().unwrap());
1036
1037 let cancelation = next_packet(&mut stub)
1038 .await
1039 .expect("dropping the future must send a cancellation");
1040 let part = cancelation.parts[0].as_ref().unwrap();
1041 assert_eq!(&part[0..4], b"rpcc");
1042 let cancel_header = proto::rpc::TRequestCancelationHeader::decode(&part[4..]).unwrap();
1043 assert_eq!(Guid::from_proto(&cancel_header.request_id), request_id);
1044 assert_eq!(cancel_header.method, "SelectRows");
1045 }
1046
1047 #[tokio::test]
1050 async fn a_completed_call_sends_no_cancellation() {
1051 let mut stub = stub_proxy(|header| {
1052 let request_id = Guid::from_proto(header.request_id.as_ref().unwrap());
1053 Some(success_reply(
1054 request_id,
1055 &proto::api::TRspPingTransaction::default(),
1056 ))
1057 })
1058 .await;
1059
1060 let connection = Connection::connect(&stub.address, None).await.unwrap();
1061 let request = proto::api::TReqPingTransaction {
1062 transaction_id: Guid::random().to_proto(),
1063 ..Default::default()
1064 };
1065 tokio::time::timeout(
1066 std::time::Duration::from_secs(10),
1067 connection.invoke::<proto::api::TRspPingTransaction>(
1068 "PingTransaction",
1069 &request,
1070 Vec::new(),
1071 None,
1072 "TRspPingTransaction",
1073 ),
1074 )
1075 .await
1076 .expect("the stub answers immediately")
1077 .unwrap();
1078
1079 let _request = next_packet(&mut stub).await.expect("the request");
1080 tokio::time::sleep(std::time::Duration::from_millis(50)).await;
1081 assert!(
1082 stub.seen.try_recv().is_err(),
1083 "a completed call must not be followed by a cancellation"
1084 );
1085 }
1086
1087 #[tokio::test]
1098 async fn the_guard_cancels_only_what_it_actually_sent() {
1099 async fn drain(receiver: &mut mpsc::Receiver<Cancellation>) -> Vec<Packet> {
1100 tokio::time::sleep(std::time::Duration::from_millis(50)).await;
1102 let mut packets = Vec::new();
1103 while let Ok(cancellation) = receiver.try_recv() {
1104 packets.push(cancellation.packet);
1105 }
1106 packets
1107 }
1108
1109 fn guard(
1110 pending: &Pending,
1111 cancels: &mpsc::Sender<Cancellation>,
1112 in_flight: &InFlight,
1113 request_id: Guid,
1114 sent: bool,
1115 ) -> PendingGuard {
1116 PendingGuard {
1117 pending: Arc::clone(pending),
1118 cancels: cancels.clone(),
1119 request_id,
1120 service: rpc::API_SERVICE.to_owned(),
1121 method: "LookupRows".to_owned(),
1122 completed: false,
1123 sent,
1124 permit: Some(
1125 Arc::clone(in_flight)
1126 .try_acquire_owned()
1127 .expect("the test never holds more than one permit"),
1128 ),
1129 }
1130 }
1131
1132 let pending: Pending = Arc::default();
1133 let in_flight = Arc::new(Semaphore::new(1));
1134 let (cancels, mut receiver) = mpsc::channel(16);
1135 let request_id = Guid::random();
1136
1137 drop(guard(&pending, &cancels, &in_flight, request_id, false));
1139 assert!(
1140 drain(&mut receiver).await.is_empty(),
1141 "cancelled a request the proxy never received"
1142 );
1143
1144 drop(guard(&pending, &cancels, &in_flight, request_id, true));
1146 let sent_packets = drain(&mut receiver).await;
1147 assert_eq!(sent_packets.len(), 1, "expected exactly one cancellation");
1148 let part = sent_packets[0].parts[0].as_ref().unwrap();
1149 assert_eq!(&part[0..4], b"rpcc");
1150 let header = proto::rpc::TRequestCancelationHeader::decode(&part[4..]).unwrap();
1151 assert_eq!(Guid::from_proto(&header.request_id), request_id);
1152
1153 let mut done = guard(&pending, &cancels, &in_flight, request_id, true);
1155 done.complete();
1156 drop(done);
1157 assert!(
1158 drain(&mut receiver).await.is_empty(),
1159 "cancelled a call that had already returned"
1160 );
1161 }
1162
1163 #[tokio::test]
1168 async fn every_in_flight_call_has_room_for_its_cancellation() {
1169 let pending: Pending = Arc::default();
1170 let in_flight = Arc::new(Semaphore::new(MAX_IN_FLIGHT));
1171 let (cancels, mut receiver) = mpsc::channel(CANCEL_QUEUE);
1172
1173 for _ in 0..MAX_IN_FLIGHT {
1174 let guard = PendingGuard {
1175 pending: Arc::clone(&pending),
1176 cancels: cancels.clone(),
1177 request_id: Guid::random(),
1178 service: rpc::API_SERVICE.to_owned(),
1179 method: "LookupRows".to_owned(),
1180 completed: false,
1181 sent: true,
1182 permit: Some(
1183 Arc::clone(&in_flight)
1184 .try_acquire_owned()
1185 .expect("the loop takes every permit exactly once"),
1186 ),
1187 };
1188 drop(guard);
1189 }
1190
1191 assert_eq!(receiver.len(), MAX_IN_FLIGHT);
1192 assert!(
1193 Arc::clone(&in_flight).try_acquire_owned().is_err(),
1194 "a queued cancellation must retain its call's permit"
1195 );
1196
1197 drop(receiver.recv().await.expect("the first cancellation"));
1200 assert!(
1201 Arc::clone(&in_flight).try_acquire_owned().is_ok(),
1202 "the consumed cancellation did not release its permit"
1203 );
1204 }
1205
1206 #[tokio::test]
1210 async fn the_in_flight_limit_bounds_pending_waiters() {
1211 let (outbound, _outbound_receiver) = mpsc::channel(MAX_IN_FLIGHT);
1212 let (cancels, _cancel_receiver) = mpsc::channel(CANCEL_QUEUE);
1213 let in_flight = Arc::new(Semaphore::new(MAX_IN_FLIGHT));
1214 let connection = Connection {
1215 outbound,
1216 cancels,
1217 pending: Arc::default(),
1218 in_flight: Arc::clone(&in_flight),
1219 address: "test".to_owned(),
1220 token: None,
1221 closed: Arc::new(AtomicBool::new(false)),
1222 reader_task: tokio::spawn(std::future::pending()),
1223 };
1224 let request = proto::api::TReqPingTransaction {
1225 transaction_id: Guid::random().to_proto(),
1226 ..Default::default()
1227 };
1228 let mut calls = Vec::with_capacity(MAX_IN_FLIGHT);
1229
1230 for _ in 0..MAX_IN_FLIGHT {
1231 let mut call = Box::pin(connection.invoke_raw(
1232 rpc::API_SERVICE,
1233 "PingTransaction",
1234 &request,
1235 Vec::new(),
1236 None,
1237 None,
1238 ));
1239 tokio::select! {
1240 biased;
1241 _ = call.as_mut() => panic!("the test connection cannot answer"),
1242 _ = tokio::task::yield_now() => {}
1243 }
1244 calls.push(call);
1245 }
1246 assert_eq!(
1247 connection.pending.lock().await.by_request.len(),
1248 MAX_IN_FLIGHT
1249 );
1250
1251 let mut overflow = Box::pin(connection.invoke_raw(
1252 rpc::API_SERVICE,
1253 "PingTransaction",
1254 &request,
1255 Vec::new(),
1256 None,
1257 None,
1258 ));
1259 tokio::select! {
1260 biased;
1261 _ = overflow.as_mut() => panic!("the overflow call cannot complete"),
1262 _ = tokio::task::yield_now() => {}
1263 }
1264 assert_eq!(
1265 connection.pending.lock().await.by_request.len(),
1266 MAX_IN_FLIGHT,
1267 "a call waiting for capacity must not register another waiter"
1268 );
1269 assert!(
1270 Arc::clone(&in_flight).try_acquire_owned().is_err(),
1271 "all in-flight permits should be held by the registered calls"
1272 );
1273
1274 drop(overflow);
1275 drop(calls);
1276 }
1277
1278 #[tokio::test]
1289 async fn a_call_that_times_out_while_queuing_cancels_nothing() {
1290 let (outbound, _outbound_receiver) = mpsc::channel(1);
1291 let (cancels, mut cancel_receiver) = mpsc::channel(CANCEL_QUEUE);
1292 let connection = Connection {
1293 outbound,
1294 cancels,
1295 pending: Arc::default(),
1296 in_flight: Arc::new(Semaphore::new(MAX_IN_FLIGHT)),
1297 address: "test".to_owned(),
1298 token: None,
1299 closed: Arc::new(AtomicBool::new(false)),
1300 reader_task: tokio::spawn(std::future::pending()),
1302 };
1303
1304 connection
1305 .outbound
1306 .try_send(Packet::message(
1307 Guid::random(),
1308 vec![Some(Bytes::from_static(b"blocker"))],
1309 PacketFlags::NONE,
1310 ))
1311 .expect("the queue starts empty");
1312
1313 let request = proto::api::TReqPingTransaction {
1314 transaction_id: Guid::random().to_proto(),
1315 ..Default::default()
1316 };
1317 let error = tokio::time::timeout(
1318 std::time::Duration::from_secs(10),
1319 connection.invoke_raw(
1320 rpc::API_SERVICE,
1321 "PingTransaction",
1322 &request,
1323 Vec::new(),
1324 Some(std::time::Duration::from_millis(50)),
1325 None,
1326 ),
1327 )
1328 .await
1329 .expect("the deadline must end a call that cannot even be queued")
1330 .unwrap_err();
1331 assert!(matches!(error, Error::Timeout { .. }), "got {error}");
1332
1333 tokio::time::sleep(std::time::Duration::from_millis(50)).await;
1335 assert!(
1336 cancel_receiver.try_recv().is_err(),
1337 "cancelled a request that never left the queue"
1338 );
1339 assert!(
1340 connection.pending.lock().await.by_request.is_empty(),
1341 "the timed-out call left its entry behind"
1342 );
1343 }
1344
1345 #[test]
1351 fn dropping_a_call_outside_a_runtime_does_not_panic() {
1352 let runtime = tokio::runtime::Builder::new_current_thread()
1353 .enable_all()
1354 .build()
1355 .unwrap();
1356
1357 let (stub, connection) = runtime.block_on(async {
1360 let stub = stub_proxy(|_| None).await;
1361 let connection = Connection::connect(&stub.address, None).await.unwrap();
1362 (stub, connection)
1363 });
1364 let request = proto::api::TReqPingTransaction {
1365 transaction_id: Guid::random().to_proto(),
1366 ..Default::default()
1367 };
1368
1369 let mut call = Box::pin(connection.invoke_raw(
1370 rpc::API_SERVICE,
1371 "PingTransaction",
1372 &request,
1373 Vec::new(),
1374 None,
1375 None,
1376 ));
1377 runtime.block_on(async {
1379 let _ = tokio::time::timeout(std::time::Duration::from_millis(50), &mut call).await;
1380 });
1381
1382 drop(call);
1385 drop(connection);
1386 drop(stub);
1387 }
1388
1389 #[tokio::test]
1390 async fn the_pending_map_does_not_leak_when_a_call_is_dropped() {
1391 let stub = stub_proxy(|_| None).await;
1392 let connection = Connection::connect(&stub.address, None).await.unwrap();
1393
1394 let request = proto::api::TReqPingTransaction {
1395 transaction_id: Guid::random().to_proto(),
1396 ..Default::default()
1397 };
1398 {
1399 let call = connection.invoke::<proto::api::TRspPingTransaction>(
1400 "PingTransaction",
1401 &request,
1402 Vec::new(),
1403 None,
1404 "TRspPingTransaction",
1405 );
1406 let _ = tokio::time::timeout(std::time::Duration::from_millis(50), call).await;
1408 }
1409
1410 tokio::time::sleep(std::time::Duration::from_millis(50)).await;
1411 assert!(
1412 connection.pending.lock().await.by_request.is_empty(),
1413 "a dropped call left its entry in the pending map"
1414 );
1415 }
1416}