1use crate::actions::{EvaluationContext, RequestRule, ResponseRule};
2use crate::errors::{DoorkeeperError, ProxyError, WorkerError};
3use crate::frame::{
4 self, FrameOpcode, FrameParams, RequestFrame, ResponseFrame, ResponseOpcode,
5 read_response_frame, write_frame,
6};
7use crate::{RequestOpcode, TargetShard};
8use bytes::Bytes;
9use compression::no_compression;
10use itertools::Either;
11use scylla_cql::frame::types::read_string_multimap;
12use std::collections::HashMap;
13use std::env::VarError;
14use std::fmt::Display;
15use std::future::Future;
16use std::net::{IpAddr, Ipv4Addr, SocketAddr};
17use std::sync::atomic::{AtomicBool, AtomicU8, AtomicU32, Ordering};
18use std::sync::{Arc, Mutex};
19use tokio::io::{AsyncRead, AsyncWrite, DuplexStream};
20use tokio::net::{TcpListener, TcpSocket, TcpStream};
21use tokio::sync::mpsc::error::TryRecvError;
22use tokio::sync::mpsc::{UnboundedReceiver, UnboundedSender};
23use tokio::sync::{broadcast, mpsc};
24use tracing::{debug, error, info, trace, warn};
25
26type FinishWaiter = mpsc::Receiver<()>;
28type FinishGuard = mpsc::Sender<()>;
29
30type TerminateNotifier = tokio::sync::broadcast::Receiver<()>;
32type TerminateSignaler = tokio::sync::broadcast::Sender<()>;
33
34type ConnectionCloseNotifier = tokio::sync::broadcast::Receiver<()>;
37type ConnectionCloseSignaler = tokio::sync::broadcast::Sender<()>;
38
39type ErrorPropagator = mpsc::UnboundedSender<ProxyError>;
42type ErrorSink = mpsc::UnboundedReceiver<ProxyError>;
43
44static HARDCODED_OPTIONS_PARAMS: FrameParams = FrameParams {
45 flags: 0,
46 version: 0x04,
47 stream: 0,
48};
49
50#[derive(Clone, Copy, Debug)]
52pub enum ShardAwareness {
53 Unaware,
55 QueryNode,
61 FixedNum(u16),
63}
64
65impl ShardAwareness {
66 pub fn is_aware(&self) -> bool {
67 !matches!(self, Self::Unaware)
68 }
69}
70
71enum NodeType {
91 Real {
92 real_addr: SocketAddr,
93 shard_awareness: ShardAwareness,
94 response_rules: Option<Vec<ResponseRule>>,
95 },
96 Simulated,
97}
98
99pub struct Node {
100 proxy_addr: SocketAddr,
101 request_rules: Option<Vec<RequestRule>>,
102 node_type: NodeType,
103}
104
105impl Node {
106 pub fn new(
108 real_addr: SocketAddr,
109 proxy_addr: SocketAddr,
110 shard_awareness: ShardAwareness,
111 request_rules: Option<Vec<RequestRule>>,
112 response_rules: Option<Vec<ResponseRule>>,
113 ) -> Self {
114 Self {
115 proxy_addr,
116 request_rules,
117 node_type: NodeType::Real {
118 real_addr,
119 shard_awareness,
120 response_rules,
121 },
122 }
123 }
124
125 pub fn new_dry_mode(proxy_addr: SocketAddr, request_rules: Option<Vec<RequestRule>>) -> Self {
127 Self {
128 proxy_addr,
129 request_rules,
130 node_type: NodeType::Simulated,
131 }
132 }
133
134 pub fn builder() -> NodeBuilder {
135 NodeBuilder {
136 real_addr: None,
137 proxy_addr: None,
138 shard_awareness: None,
139 request_rules: None,
140 response_rules: None,
141 }
142 }
143}
144
145pub struct NodeBuilder {
146 real_addr: Option<SocketAddr>,
147 proxy_addr: Option<SocketAddr>,
148 shard_awareness: Option<ShardAwareness>,
149 request_rules: Option<Vec<RequestRule>>,
150 response_rules: Option<Vec<ResponseRule>>,
151}
152
153impl NodeBuilder {
154 pub fn real_address(mut self, real_addr: SocketAddr) -> Self {
155 self.real_addr = Some(real_addr);
156 self
157 }
158
159 pub fn proxy_address(mut self, proxy_addr: SocketAddr) -> Self {
160 self.proxy_addr = Some(proxy_addr);
161 self
162 }
163
164 pub fn shard_awareness(mut self, shard_awareness: ShardAwareness) -> Self {
165 self.shard_awareness = Some(shard_awareness);
166 self
167 }
168
169 pub fn request_rules(mut self, request_rules: Vec<RequestRule>) -> Self {
170 self.request_rules = Some(request_rules);
171 self
172 }
173
174 pub fn response_rules(mut self, response_rules: Vec<ResponseRule>) -> Self {
175 self.response_rules = Some(response_rules);
176 self
177 }
178
179 pub fn build(self) -> Node {
181 Node {
182 proxy_addr: self.proxy_addr.expect("Proxy addr is required!"),
183 request_rules: self.request_rules,
184 node_type: NodeType::Real {
185 real_addr: self.real_addr.expect("Real addr is required!"),
186 shard_awareness: self.shard_awareness.expect("Shard awareness is required!"),
187 response_rules: self.response_rules,
188 },
189 }
190 }
191
192 pub fn build_dry_mode(self) -> Node {
194 Node {
195 proxy_addr: self.proxy_addr.expect("Proxy addr is required!"),
196 request_rules: self.request_rules,
197 node_type: NodeType::Simulated,
198 }
199 }
200}
201
202#[derive(Clone, Copy)]
203struct DisplayableRealAddrOption(Option<SocketAddr>);
204impl Display for DisplayableRealAddrOption {
205 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
206 if let Some(addr) = self.0 {
207 write!(f, "{addr}")
208 } else {
209 write!(f, "<dry mode>")
210 }
211 }
212}
213
214#[derive(Clone, Copy)]
215struct DisplayableShard(Option<TargetShard>);
216impl Display for DisplayableShard {
217 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
218 if let Some(shard) = self.0 {
219 write!(f, "shard {shard}")
220 } else {
221 write!(f, "unknown shard")
222 }
223 }
224}
225
226enum InternalNode {
227 Real {
228 real_addr: SocketAddr,
229 proxy_addr: SocketAddr,
230 shard_awareness: ShardAwareness,
231 request_rules: Arc<Mutex<Vec<RequestRule>>>,
232 response_rules: Arc<Mutex<Vec<ResponseRule>>>,
233 },
234 Simulated {
235 proxy_addr: SocketAddr,
236 request_rules: Arc<Mutex<Vec<RequestRule>>>,
237 },
238}
239
240impl InternalNode {
241 fn proxy_addr(&self) -> SocketAddr {
242 match *self {
243 InternalNode::Real { proxy_addr, .. } => proxy_addr,
244 InternalNode::Simulated { proxy_addr, .. } => proxy_addr,
245 }
246 }
247 fn real_addr(&self) -> Option<SocketAddr> {
248 match *self {
249 InternalNode::Real { real_addr, .. } => Some(real_addr),
250 InternalNode::Simulated { .. } => None,
251 }
252 }
253 fn request_rules(&self) -> &Arc<Mutex<Vec<RequestRule>>> {
254 match self {
255 InternalNode::Real { request_rules, .. } => request_rules,
256 InternalNode::Simulated { request_rules, .. } => request_rules,
257 }
258 }
259}
260
261impl From<Node> for InternalNode {
262 fn from(node: Node) -> Self {
263 match node.node_type {
264 NodeType::Real {
265 real_addr,
266 shard_awareness,
267 response_rules,
268 } => InternalNode::Real {
269 real_addr,
270 proxy_addr: node.proxy_addr,
271 shard_awareness,
272 request_rules: node
273 .request_rules
274 .map(|rules| Arc::new(Mutex::new(rules)))
275 .unwrap_or_default(),
276 response_rules: response_rules
277 .map(|rules| Arc::new(Mutex::new(rules)))
278 .unwrap_or_default(),
279 },
280 NodeType::Simulated => InternalNode::Simulated {
281 proxy_addr: node.proxy_addr,
282 request_rules: node
283 .request_rules
284 .map(|rules| Arc::new(Mutex::new(rules)))
285 .unwrap_or_default(),
286 },
287 }
288 }
289}
290
291pub struct ProxyBuilder {
292 nodes: Vec<Node>,
293}
294
295impl ProxyBuilder {
296 pub fn with_node(mut self, node: Node) -> ProxyBuilder {
297 self.nodes.push(node);
298 self
299 }
300
301 pub fn build(self) -> Proxy {
302 Proxy::new(self.nodes)
303 }
304}
305
306pub struct Proxy {
307 nodes: Vec<InternalNode>,
308}
309
310impl Proxy {
311 pub fn new(nodes: impl IntoIterator<Item = Node>) -> Self {
312 Proxy {
313 nodes: nodes.into_iter().map(|node| node.into()).collect(),
314 }
315 }
316
317 pub fn builder() -> ProxyBuilder {
318 ProxyBuilder { nodes: vec![] }
319 }
320
321 pub fn translation_map(&self) -> HashMap<SocketAddr, SocketAddr> {
325 let mut translation_map = HashMap::new();
326 for node in self.nodes.iter() {
327 if let &InternalNode::Real {
328 real_addr,
329 proxy_addr,
330 ..
331 } = node
332 {
333 translation_map.insert(real_addr, proxy_addr);
334 let shard_aware_real_addr = SocketAddr::new(real_addr.ip(), 19042);
335 translation_map.insert(shard_aware_real_addr, proxy_addr);
336 }
337 }
338 translation_map
339 }
340
341 pub async fn run(self) -> Result<RunningProxy, DoorkeeperError> {
344 let (terminate_signaler, _t) = tokio::sync::broadcast::channel(1);
345 let (finish_guard, finish_waiter) = mpsc::channel(1);
346
347 let (error_propagator, error_sink) = mpsc::unbounded_channel();
348 let (addresses_and_doorkeepers, running_nodes): (Vec<_>, Vec<RunningNode>) = self
349 .nodes
350 .into_iter()
351 .map(|node| {
352 let cc_event_sender = Arc::new(Mutex::new(HashMap::new()));
353 let running = {
354 let (request_rules, response_rules) = match node {
355 InternalNode::Real {
356 ref request_rules,
357 ref response_rules,
358 ..
359 } => (request_rules, Some(response_rules)),
360 InternalNode::Simulated {
361 ref request_rules, ..
362 } => (request_rules, None),
363 };
364 RunningNode {
365 request_rules: request_rules.clone(),
366 response_rules: response_rules.cloned(),
367 cc_event_sender: cc_event_sender.clone(),
368 }
369 };
370 let proxy_addr = node.proxy_addr();
371 let doorkeeper = {
372 Doorkeeper::spawn(
373 node,
374 terminate_signaler.clone(),
375 finish_guard.clone(),
376 error_propagator.clone(),
377 cc_event_sender,
378 )
379 };
380 let address_and_doorkeeper = async move { Ok((proxy_addr, doorkeeper.await?)) };
381 (address_and_doorkeeper, running)
382 })
383 .unzip();
384
385 let mut duplex_submitter_map = HashMap::new();
386 for address_and_doorkeeper in addresses_and_doorkeepers {
387 let (proxy_addr, duplex_submitter) = address_and_doorkeeper.await?; duplex_submitter_map.insert(proxy_addr, duplex_submitter);
389 }
390
391 let transport_factory = Arc::new(TransportFactory {
392 duplex_submitter_map,
393 });
394
395 Ok(RunningProxy {
396 terminate_signaler,
397 finish_waiter,
398 running_nodes,
399 error_sink,
400 transport_factory,
401 })
402 }
403}
404
405pub struct TransportFactory {
418 duplex_submitter_map: HashMap<SocketAddr, UnboundedSender<(DuplexStream, SocketAddr)>>,
419}
420
421impl TransportFactory {
422 pub fn connect(
423 &self,
424 connect_address: SocketAddr,
425 source_ip: Option<IpAddr>,
426 source_port: Option<u16>,
427 ) -> Result<tokio::io::DuplexStream, std::io::Error> {
428 let sender = self
429 .duplex_submitter_map
430 .get(&connect_address)
431 .ok_or_else(|| {
432 std::io::Error::new(
433 std::io::ErrorKind::HostUnreachable,
434 format!("No host under address {connect_address}"),
435 )
436 })?;
437
438 let source_ip = source_ip.unwrap_or(IpAddr::V4(Ipv4Addr::LOCALHOST));
439 let source_port = source_port.unwrap_or_else(|| rand::random_range(49152..=65535));
440 let source_address = SocketAddr::from((source_ip, source_port));
441
442 let (duplex1, duplex2) = tokio::io::duplex(8 * 1024);
443 sender.send((duplex2, source_address)).map_err(|_| {
444 std::io::Error::new(
445 std::io::ErrorKind::ConnectionRefused,
446 format!("The proxy node {connect_address} no longer accepts connections"),
447 )
448 })?;
449 Ok(duplex1)
450 }
451}
452
453pub struct RunningNode {
455 request_rules: Arc<Mutex<Vec<RequestRule>>>,
456 response_rules: Option<Arc<Mutex<Vec<ResponseRule>>>>,
457
458 cc_event_sender: Arc<Mutex<HashMap<usize, mpsc::UnboundedSender<ResponseFrame>>>>,
468}
469
470impl RunningNode {
471 pub fn change_request_rules(&mut self, rules: Option<Vec<RequestRule>>) {
473 *self.request_rules.lock().unwrap() = rules.unwrap_or_default();
474 }
475
476 pub fn append_request_rules(&mut self, mut rules: Vec<RequestRule>) {
478 self.request_rules.lock().unwrap().append(&mut rules);
479 }
480
481 pub fn prepend_request_rules(&mut self, rules: Vec<RequestRule>) {
483 let mut new_rules = rules;
484 let mut old_rules_guard = self.request_rules.lock().unwrap();
485 new_rules.append(&mut *old_rules_guard);
486 *old_rules_guard = new_rules;
487 }
488
489 pub fn change_response_rules(&mut self, rules: Option<Vec<ResponseRule>>) {
491 *self
492 .response_rules
493 .as_ref()
494 .expect("No response rules on a simulated node!")
495 .lock()
496 .unwrap() = rules.unwrap_or_default();
497 }
498
499 pub fn append_response_rules(&mut self, mut rules: Vec<ResponseRule>) {
501 self.response_rules
502 .as_ref()
503 .expect("No response rules on a simulated node!")
504 .lock()
505 .unwrap()
506 .append(&mut rules);
507 }
508
509 pub fn prepend_response_rules(&mut self, rules: Vec<ResponseRule>) {
511 let mut old_rules_guard = self
512 .response_rules
513 .as_ref()
514 .expect("No response rules on a simulated node!")
515 .lock()
516 .unwrap();
517 let mut new_rules = rules;
518 new_rules.append(&mut *old_rules_guard);
519 *old_rules_guard = new_rules;
520 }
521
522 pub fn inject_event_to_cc(&self, body: Bytes) -> bool {
532 let mut guard = self.cc_event_sender.lock().unwrap();
533 if guard.is_empty() {
534 return false;
535 }
536 let mut any_sent = false;
537 guard.retain(|_conn_no, tx| {
538 let frame = ResponseFrame::new(
539 FrameParams {
540 version: 4,
541 flags: 0,
542 stream: -1,
543 }
544 .for_response(),
545 ResponseOpcode::Event,
546 body.clone(),
547 );
548 let ok = tx.send(frame).is_ok();
549 any_sent |= ok;
550 ok
553 });
554 any_sent
555 }
556}
557
558pub struct RunningProxy {
560 terminate_signaler: TerminateSignaler,
561 finish_waiter: FinishWaiter,
562 pub running_nodes: Vec<RunningNode>,
563 error_sink: ErrorSink,
564 transport_factory: Arc<TransportFactory>,
565}
566
567impl RunningProxy {
568 pub fn turn_off_rules(&mut self) {
570 for (request_rules, response_rules) in self
571 .running_nodes
572 .iter_mut()
573 .map(|node| (&node.request_rules, &node.response_rules))
574 {
575 request_rules.lock().unwrap().clear();
576 if let Some(response_rules) = response_rules {
577 response_rules.lock().unwrap().clear();
578 }
579 }
580 }
581
582 pub fn sanity_check(&mut self) -> Result<(), ProxyError> {
585 match self.error_sink.try_recv() {
586 Ok(err) => Err(err),
587 Err(TryRecvError::Empty) => Ok(()),
588 Err(TryRecvError::Disconnected) => {
589 Err(ProxyError::SanityCheckFailure)
591 }
592 }
593 }
594
595 pub async fn wait_for_error(&mut self) -> Option<ProxyError> {
597 self.error_sink.recv().await
598 }
599
600 pub fn transport_factory(&self) -> Arc<TransportFactory> {
602 Arc::clone(&self.transport_factory)
603 }
604
605 pub async fn finish(mut self) -> Result<(), ProxyError> {
608 self.terminate_signaler.send(()).map_err(|err| {
609 ProxyError::AwaitFinishFailure(format!(
610 "Send error in terminate_signaler: {err} (bug!)"
611 ))
612 })?;
613 info!("Sent finish signal to proxy workers.");
614
615 std::mem::drop(self.terminate_signaler);
617
618 if self.finish_waiter.recv().await.is_some() {
619 unreachable!();
620 };
621 info!("All workers have finished.");
622
623 match self.error_sink.try_recv() {
624 Ok(err) => Err(err),
625 Err(TryRecvError::Disconnected) => Ok(()),
626 Err(TryRecvError::Empty) => {
627 unreachable!("Worker await logic bug!");
629 }
630 }
631 }
632}
633
634struct Doorkeeper {
639 node: InternalNode,
640 listener: TcpListener,
641 terminate_signaler: TerminateSignaler,
642 finish_guard: FinishGuard,
643 shards_count: Option<u16>,
644 error_propagator: ErrorPropagator,
645 cc_event_sender: Arc<Mutex<HashMap<usize, mpsc::UnboundedSender<ResponseFrame>>>>,
646 duplex_receiver: UnboundedReceiver<(DuplexStream, SocketAddr)>,
647 _duplex_submitter: UnboundedSender<(DuplexStream, SocketAddr)>,
651}
652
653impl Doorkeeper {
654 async fn spawn(
655 node: InternalNode,
656 terminate_signaler: TerminateSignaler,
657 finish_guard: FinishGuard,
658 error_propagator: ErrorPropagator,
659 cc_event_sender: Arc<Mutex<HashMap<usize, mpsc::UnboundedSender<ResponseFrame>>>>,
660 ) -> Result<UnboundedSender<(DuplexStream, SocketAddr)>, DoorkeeperError> {
661 let listener = TcpListener::bind(node.proxy_addr())
662 .await
663 .map_err(|err| DoorkeeperError::DriverConnectionAttempt(node.proxy_addr(), err))?;
664
665 if let InternalNode::Real {
666 shard_awareness,
667 real_addr,
668 ..
669 } = node
670 {
671 info!(
672 "Spawned a {} doorkeeper for pair real:{} - proxy:{}.",
673 if shard_awareness.is_aware() {
674 "shard-aware"
675 } else {
676 "shard-unaware"
677 },
678 real_addr,
679 node.proxy_addr(),
680 );
681 } else {
682 info!(
683 "Spawned a dry-mode doorkeeper for proxy:{}.",
684 node.proxy_addr(),
685 )
686 };
687
688 let (duplex_submitter, duplex_receiver) = mpsc::unbounded_channel();
689 let doorkeeper = Doorkeeper {
690 shards_count: None, node,
692 listener,
693 terminate_signaler,
694 finish_guard,
695 error_propagator,
696 cc_event_sender,
697 duplex_receiver,
698 _duplex_submitter: duplex_submitter.clone(),
699 };
700 tokio::task::spawn(doorkeeper.run());
701 Ok(duplex_submitter)
702 }
703
704 async fn run(mut self) {
705 self.update_shards_count().await;
706 let mut own_terminate_notifier = self.terminate_signaler.subscribe();
707 let (connection_close_tx, _connection_close_rx) = broadcast::channel::<()>(2);
708 let mut connection_no: usize = 0;
709 loop {
710 tokio::select! {
711 res = self.accept_connection(&connection_close_tx, connection_no) => {
712 match res {
713 Ok(()) => connection_no += 1,
714 Err(err) => {
715 error!(
716 "Error in doorkeeper with addr {} for node {}: {}",
717 self.node.proxy_addr(),
718 DisplayableRealAddrOption(self.node.real_addr()),
719 err
720 );
721 let _ = self.error_propagator.send(err.into());
722 break;
723 },
724 }
725 },
726 _terminate = own_terminate_notifier.recv() => break
727 }
728 }
729 debug!(
730 "Doorkeeper exits: proxy {}, node {}.",
731 self.node.proxy_addr(),
732 DisplayableRealAddrOption(self.node.real_addr())
733 );
734 }
735
736 async fn update_shards_count(&mut self) {
737 if let InternalNode::Real {
738 real_addr,
739 shard_awareness,
740 ..
741 } = self.node
742 {
743 self.shards_count = match shard_awareness {
744 ShardAwareness::Unaware => None,
745 ShardAwareness::FixedNum(shards_num) => Some(shards_num),
746 ShardAwareness::QueryNode => match self.obtain_shards_count(real_addr).await {
747 Ok(shards) => Some(shards),
748 Err(DoorkeeperError::ObtainingShardNumberNoShardInfo) => {
750 info!(
751 "Doorkeeper with addr {} found no shard info in node {}; falling back to ShardAwareness::Unaware",
752 self.node.proxy_addr(),
753 DisplayableRealAddrOption(self.node.real_addr()),
754 );
755 None
756 }
757 Err(e) => {
758 error!(
759 "Error in doorkeeper with addr {} while querying shard info from node {}: {}",
760 self.node.proxy_addr(),
761 DisplayableRealAddrOption(self.node.real_addr()),
762 e
763 );
764 None
765 }
766 },
767 }
768 }
769 }
770
771 #[expect(clippy::too_many_arguments)]
772 async fn spawn_workers(
773 &mut self,
774 driver_addr: SocketAddr,
775 connection_close_tx: &ConnectionCloseSignaler,
776 connection_no: usize,
777 driver_read: impl AsyncRead + Send + Unpin + 'static,
778 driver_write: impl AsyncWrite + Send + Unpin + 'static,
779 cluster_stream: Option<TcpStream>,
780 shard: Option<TargetShard>,
781 ) {
782 let new_worker = || ProxyWorker {
783 terminate_notifier: self.terminate_signaler.subscribe(),
784 finish_guard: self.finish_guard.clone(),
785 connection_close_notifier: connection_close_tx.subscribe(),
786 error_propagator: self.error_propagator.clone(),
787 driver_addr,
788 real_addr: self.node.real_addr(),
789 proxy_addr: self.node.proxy_addr(),
790 shard,
791 };
792
793 let (tx_request, rx_request) = mpsc::unbounded_channel::<RequestFrame>();
794 let (tx_response, rx_response) = mpsc::unbounded_channel::<ResponseFrame>();
795 let (tx_cluster, rx_cluster) = mpsc::unbounded_channel::<RequestFrame>();
796 let (tx_driver, rx_driver) = mpsc::unbounded_channel::<ResponseFrame>();
797 let event_register_flag = Arc::new(AtomicBool::new(false));
798
799 let (
800 compression_writer_request_processor,
801 compression_reader_receiver_from_driver,
802 compression_reader_receiver_from_cluster,
803 compression_reader_sender_to_driver,
804 compression_reader_sender_to_cluster,
805 ) = compression::make_compression_infra();
806
807 {
808 let worker = new_worker();
809 tokio::task::spawn(async move {
810 worker
811 .receiver_from_driver(
812 driver_read,
813 tx_request,
814 compression_reader_receiver_from_driver,
815 )
816 .await;
817 });
818 }
819 {
820 let worker = new_worker();
821 let conn_close_sub = connection_close_tx.subscribe();
822 let term_sub = self.terminate_signaler.subscribe();
823 tokio::task::spawn(async move {
824 worker
825 .sender_to_driver(
826 driver_write,
827 rx_driver,
828 conn_close_sub,
829 term_sub,
830 compression_reader_sender_to_driver,
831 )
832 .await;
833 });
834 }
835 {
836 let worker = new_worker();
837 let request_rules = Arc::clone(self.node.request_rules());
838 let conn_close = connection_close_tx.clone();
839 let event_flag = Arc::clone(&event_register_flag);
840 let tx_driver_clone = tx_driver.clone();
841 let tx_cluster_clone = tx_cluster.clone();
842 let cc_sender = Arc::clone(&self.cc_event_sender);
843 tokio::task::spawn(async move {
844 worker
845 .request_processor(
846 rx_request,
847 tx_driver_clone,
848 tx_cluster_clone,
849 connection_no,
850 request_rules,
851 conn_close,
852 event_flag,
853 compression_writer_request_processor,
854 cc_sender,
855 )
856 .await;
857 });
858 }
859 if let InternalNode::Real {
860 ref response_rules, ..
861 } = self.node
862 {
863 let (cluster_read, cluster_write) = cluster_stream.unwrap().into_split();
864 {
865 let worker = new_worker();
866 let conn_close_sub = connection_close_tx.subscribe();
867 let term_sub = self.terminate_signaler.subscribe();
868 tokio::task::spawn(async move {
869 worker
870 .sender_to_cluster(
871 cluster_write,
872 rx_cluster,
873 conn_close_sub,
874 term_sub,
875 compression_reader_sender_to_cluster,
876 )
877 .await;
878 });
879 }
880 {
881 let worker = new_worker();
882 tokio::task::spawn(async move {
883 worker
884 .receiver_from_cluster(
885 cluster_read,
886 tx_response,
887 compression_reader_receiver_from_cluster,
888 )
889 .await;
890 });
891 }
892 {
893 let worker = new_worker();
894 let response_rules = Arc::clone(response_rules);
895 let conn_close = connection_close_tx.clone();
896 let event_flag = Arc::clone(&event_register_flag);
897 tokio::task::spawn(async move {
898 worker
899 .response_processor(
900 rx_response,
901 tx_driver,
902 tx_cluster,
903 connection_no,
904 response_rules,
905 conn_close,
906 event_flag,
907 )
908 .await;
909 });
910 }
911 }
912 debug!(
913 "Doorkeeper with addr {} of node {} spawned workers.",
914 self.node.proxy_addr(),
915 DisplayableRealAddrOption(self.node.real_addr())
916 );
917 }
918
919 async fn accept_connection(
920 &mut self,
921 connection_close_tx: &ConnectionCloseSignaler,
922 connection_no: usize,
923 ) -> Result<(), DoorkeeperError> {
924 let (driver_stream, driver_addr) = self.make_driver_stream(connection_no).await?;
925 let (cluster_stream, shard) = match self.node {
926 InternalNode::Real { real_addr, .. } => {
927 let (cluster_stream, shard) =
928 self.make_cluster_stream(driver_addr, real_addr).await?;
929 (Some(cluster_stream), shard)
930 }
931 InternalNode::Simulated { .. } => (None, None),
932 };
933
934 let (driver_read, driver_write): (
935 Box<dyn AsyncRead + Send + Unpin + 'static>,
936 Box<dyn AsyncWrite + Send + Unpin + 'static>,
937 ) = driver_stream
938 .map_either(TcpStream::into_split, tokio::io::split)
939 .either(
940 |(l, r)| (Box::new(l) as _, Box::new(r) as _),
941 |(l, r)| (Box::new(l) as _, Box::new(r) as _),
942 );
943
944 self.spawn_workers(
945 driver_addr,
946 connection_close_tx,
947 connection_no,
948 driver_read,
949 driver_write,
950 cluster_stream,
951 shard,
952 )
953 .await;
954
955 Ok(())
956 }
957
958 async fn make_driver_stream(
959 &mut self,
960 connection_no: usize,
961 ) -> Result<(Either<TcpStream, DuplexStream>, SocketAddr), DoorkeeperError> {
962 let (driver_stream, driver_addr) = tokio::select! {
963 v = self.listener.accept() => {
964 let (stream, addr) = v.map_err(|err| DoorkeeperError::DriverConnectionAttempt(self.node.proxy_addr(), err))?;
965 (Either::Left(stream), addr)
966 },
967 v = self.duplex_receiver.recv(), if !self.duplex_receiver.is_closed() => {
968 let (duplex, addr) = v.expect("receiver should be open because Doorkeeper keeps the sender alive");
973 (Either::Right(duplex), addr)
974 },
975 };
976 info!(
977 "Connected driver from {} to {}, connection no={}.",
978 driver_addr,
979 self.node.proxy_addr(),
980 connection_no
981 );
982 Ok((driver_stream, driver_addr))
983 }
984
985 async fn make_cluster_stream(
986 &mut self,
987 driver_addr: SocketAddr,
988 real_addr: SocketAddr,
989 ) -> Result<(TcpStream, Option<TargetShard>), DoorkeeperError> {
990 let mut cluster_stream = if let Some(shards) = self.shards_count {
991 let (socket, mut desired_addr) =
992 open_shard_aware_socket_to_real_node(real_addr, driver_addr.port())
993 .map_err(DoorkeeperError::SocketCreate)?;
994
995 let shard_preserving_addr = {
996 while socket.bind(desired_addr).is_err() {
997 let next_port = self.next_port_to_same_shard(desired_addr.port());
999 if next_port == driver_addr.port() {
1000 return Err(DoorkeeperError::NoMorePorts);
1001 }
1002 desired_addr.set_port(next_port);
1003 }
1004 desired_addr
1005 };
1006
1007 let stream = socket.connect(real_addr).await;
1008 if let Ok(ok) = &stream {
1009 info!(
1010 "Connected to the cluster from {} at {}, intended shard {}.",
1011 ok.local_addr().unwrap(),
1012 real_addr,
1013 shard_preserving_addr.port() % shards
1014 );
1015 }
1016 stream
1017 } else {
1018 let stream = TcpStream::connect(real_addr).await;
1019 if stream.is_ok() {
1020 info!("Connected to the cluster at {}.", real_addr);
1021 }
1022 stream
1023 }
1024 .map_err(|err| DoorkeeperError::NodeConnectionAttempt(real_addr, err))?;
1025
1026 let shard = if self.shards_count.is_some() {
1033 self.obtain_shard_number(real_addr, &mut cluster_stream)
1034 .await?
1035 } else {
1036 None
1037 };
1038
1039 Ok((cluster_stream, shard))
1040 }
1041
1042 fn next_port_to_same_shard(&self, port: u16) -> u16 {
1043 port.wrapping_add(self.shards_count.unwrap())
1044 }
1045
1046 async fn get_supported_options(
1047 connection: &mut TcpStream,
1048 ) -> Result<HashMap<String, Vec<String>>, DoorkeeperError> {
1049 write_frame(
1050 HARDCODED_OPTIONS_PARAMS,
1051 FrameOpcode::Request(RequestOpcode::Options),
1052 &Bytes::new(),
1053 connection,
1054 &no_compression(),
1055 )
1056 .await
1057 .map_err(DoorkeeperError::ObtainingShardNumber)?;
1058
1059 let supported_frame = read_response_frame(connection, &compression::no_compression())
1060 .await
1061 .map_err(DoorkeeperError::ObtainingShardNumberFrame)?;
1062
1063 let options = read_string_multimap(&mut supported_frame.body.as_ref())
1064 .map_err(DoorkeeperError::ObtainingShardNumberParseOptions)?;
1065
1066 Ok(options)
1067 }
1068
1069 async fn obtain_shards_count(&self, real_addr: SocketAddr) -> Result<u16, DoorkeeperError> {
1070 let mut connection = TcpStream::connect(real_addr)
1071 .await
1072 .map_err(|err| DoorkeeperError::NodeConnectionAttempt(real_addr, err))?;
1073 let options = Self::get_supported_options(&mut connection).await?;
1074 let nr_shards_entry = options.get("SCYLLA_NR_SHARDS");
1075 let shards = match nr_shards_entry
1076 .and_then(|vec| vec.first())
1077 .ok_or(DoorkeeperError::ObtainingShardNumberNoShardInfo)?
1078 .parse::<u16>()
1079 .map_err(DoorkeeperError::ObtainingShardNumberParseShardNumber)?
1080 {
1081 0u16 => Err(DoorkeeperError::ObtainingShardNumberGotZero),
1082 num => Ok(num),
1083 }?;
1084 info!("Obtained shards number on node {}: {}", real_addr, shards);
1085 Ok(shards)
1086 }
1087
1088 async fn obtain_shard_number(
1089 &self,
1090 real_addr: SocketAddr,
1091 connection: &mut TcpStream,
1092 ) -> Result<Option<TargetShard>, DoorkeeperError> {
1093 let options = Self::get_supported_options(connection).await?;
1094 let shard_entry = options.get("SCYLLA_SHARD");
1095 let shard = shard_entry
1096 .and_then(|vec| vec.first())
1097 .map(|s| {
1098 s.parse::<u16>()
1099 .map_err(DoorkeeperError::ObtainingShardNumberParseShardNumber)
1100 })
1101 .transpose()?;
1102 info!("Connected to node {}, shard {:?}", real_addr, shard);
1103 Ok(shard)
1104 }
1105}
1106
1107mod compression {
1108 use std::error::Error;
1109 use std::sync::{Arc, OnceLock};
1110
1111 use bytes::Bytes;
1112 use scylla_cql::frame::frame_errors::{
1113 CqlRequestSerializationError, FrameBodyExtensionsParseError,
1114 };
1115 use scylla_cql::frame::request::{
1116 DeserializableRequest as _, RequestDeserializationError, Startup, options,
1117 };
1118 use scylla_cql::frame::{Compression, compress_append, decompress, flag};
1119 use tracing::{error, warn};
1120
1121 #[derive(Debug, thiserror::Error)]
1122 pub(crate) enum CompressionError {
1123 #[error("Snap compression error: {0}")]
1125 SnapCompressError(Arc<dyn Error + Sync + Send>),
1126
1127 #[error("Frame is to be compressed, but no compression negotiated for connection.")]
1129 NoCompressionNegotiated,
1130 }
1131
1132 type CompressionInfo = Arc<OnceLock<Option<Compression>>>;
1133
1134 #[derive(Debug, Clone)]
1139 pub(crate) struct CompressionWriter(CompressionInfo);
1140 impl CompressionWriter {
1141 pub(crate) fn set(
1142 &self,
1143 compression: Option<Compression>,
1144 ) -> Result<(), Option<Compression>> {
1145 self.0.set(compression)
1146 }
1147
1148 pub(crate) fn set_from_startup(
1149 &self,
1150 mut body: &[u8],
1151 ) -> Result<Option<Compression>, RequestDeserializationError> {
1152 let startup = Startup::deserialize_with_features(&mut body, &Default::default())?;
1153 let maybe_compression = startup.options.get(options::COMPRESSION);
1154 let maybe_compression = maybe_compression.and_then(|compression| {
1155 compression
1156 .parse::<Compression>()
1157 .inspect_err(|err| error!("STARTUP compression error: {}", err))
1158 .ok()
1159 });
1160 let _ = self.set(maybe_compression).inspect_err(|_| {
1161 warn!("Captured second or further STARTUP frame on the same connection")
1162 });
1163
1164 Ok(maybe_compression)
1165 }
1166 }
1167
1168 #[derive(Debug, Clone)]
1172 pub(crate) struct CompressionReader(CompressionInfo);
1173 impl CompressionReader {
1174 pub(crate) fn get(&self) -> Option<Option<Compression>> {
1179 self.0.get().copied()
1180 }
1181
1182 pub(crate) fn maybe_compress_body(
1183 &self,
1184 flags: u8,
1185 body: &[u8],
1186 ) -> Result<Option<Bytes>, CompressionError> {
1187 match (flags & flag::COMPRESSION != 0, self.get().flatten()) {
1188 (true, Some(compression)) => {
1189 let mut buf = Vec::new();
1190 compress_append(body, compression, &mut buf).map_err(|err| {
1191 let CqlRequestSerializationError::SnapCompressError(err) = err else {
1192 unreachable!("BUG: compress_append returned variant different than SnapCompressError")
1193 };
1194 CompressionError::SnapCompressError(err)
1195 })?;
1196 Ok(Some(Bytes::from(buf)))
1197 }
1198 (true, None) => Err(CompressionError::NoCompressionNegotiated),
1199 (false, _) => Ok(None),
1200 }
1201 }
1202
1203 pub(crate) fn maybe_decompress_body(
1204 &self,
1205 flags: u8,
1206 body: Bytes,
1207 ) -> Result<Bytes, FrameBodyExtensionsParseError> {
1208 match (flags & flag::COMPRESSION != 0, self.get().flatten()) {
1209 (true, Some(compression)) => decompress(&body, compression).map(Into::into),
1210 (true, None) => Err(FrameBodyExtensionsParseError::NoCompressionNegotiated),
1211 (false, _) => Ok(body),
1212 }
1213 }
1214 }
1215
1216 pub(crate) fn make_compression_infra() -> (
1217 CompressionWriter,
1218 CompressionReader,
1219 CompressionReader,
1220 CompressionReader,
1221 CompressionReader,
1222 ) {
1223 let info = Arc::new(OnceLock::new());
1224 (
1225 CompressionWriter(info.clone()),
1226 CompressionReader(info.clone()),
1227 CompressionReader(info.clone()),
1228 CompressionReader(info.clone()),
1229 CompressionReader(info),
1230 )
1231 }
1232
1233 fn mock_compression_reader(compression: Option<Compression>) -> CompressionReader {
1234 CompressionReader(Arc::new({
1235 let once = OnceLock::new();
1236 once.set(compression).unwrap();
1237 once
1238 }))
1239 }
1240
1241 pub(crate) fn no_compression() -> CompressionReader {
1243 mock_compression_reader(None)
1244 }
1245
1246 #[cfg(test)] pub(crate) fn with_compression(compression: Compression) -> CompressionReader {
1249 mock_compression_reader(Some(compression))
1250 }
1251}
1252pub(crate) use compression::{CompressionReader, CompressionWriter};
1253
1254struct ProxyWorker {
1255 terminate_notifier: TerminateNotifier,
1256 finish_guard: FinishGuard,
1257 connection_close_notifier: ConnectionCloseNotifier,
1258 error_propagator: ErrorPropagator,
1259 driver_addr: SocketAddr,
1260 real_addr: Option<SocketAddr>,
1261 proxy_addr: SocketAddr,
1262 shard: Option<TargetShard>,
1263}
1264
1265impl ProxyWorker {
1266 fn exit(self, duty: &'static str) {
1267 debug!(
1268 "Worker exits: [driver: {}, proxy: {}, node: {}, {}]::{}.",
1269 self.driver_addr,
1270 self.proxy_addr,
1271 DisplayableRealAddrOption(self.real_addr),
1272 DisplayableShard(self.shard),
1273 duty
1274 );
1275 std::mem::drop(self.finish_guard);
1276 }
1277
1278 async fn run_until_interrupted<F, Fut>(mut self, worker_name: &'static str, f: F)
1279 where
1280 F: FnOnce(SocketAddr, SocketAddr, Option<SocketAddr>) -> Fut,
1281 Fut: Future<Output = Result<(), ProxyError>>,
1282 {
1283 let fut = f(self.driver_addr, self.proxy_addr, self.real_addr);
1284
1285 tokio::select! {
1286 result = fut => {
1287 if let Err(err) = result {
1288 let _ = self.error_propagator.send(err);
1290 }
1291 }
1292 _ = self.terminate_notifier.recv() => (),
1293 _ = self.connection_close_notifier.recv() => (),
1294 }
1295 self.exit(worker_name);
1296 }
1297
1298 async fn receiver_from_driver(
1299 self,
1300 mut read_half: impl AsyncRead + Unpin,
1301 request_processor_tx: mpsc::UnboundedSender<RequestFrame>,
1302 compression: CompressionReader,
1303 ) {
1304 let shard = self.shard;
1305 self.run_until_interrupted(
1306 "receiver_from_driver",
1307 |driver_addr, proxy_addr, _real_addr| async move {
1308 loop {
1309 let frame = frame::read_request_frame(&mut read_half, &compression)
1310 .await
1311 .map_err(|err| {
1312 warn!("Request reception from {} error: {}", driver_addr, err);
1313 WorkerError::DriverDisconnected(driver_addr)
1314 })?;
1315
1316 debug!(
1317 "Intercepted Driver ({}) -> Cluster ({}) ({}) frame. opcode: {:?}.",
1318 driver_addr,
1319 proxy_addr,
1320 DisplayableShard(shard),
1321 &frame.opcode
1322 );
1323 if request_processor_tx.send(frame).is_err() {
1324 warn!("request_processor had exited.");
1325 return Result::<(), ProxyError>::Ok(());
1326 }
1327 }
1328 },
1329 )
1330 .await
1331 }
1332
1333 async fn receiver_from_cluster(
1334 self,
1335 mut read_half: impl AsyncRead + Unpin,
1336 response_processor_tx: mpsc::UnboundedSender<ResponseFrame>,
1337 compression: CompressionReader,
1338 ) {
1339 let shard = self.shard;
1340 self.run_until_interrupted(
1341 "receiver_from_cluster",
1342 |driver_addr, _proxy_addr, real_addr| async move {
1343 let real_addr = real_addr.expect("BUG: no real_addr in cluster worker");
1344 loop {
1345 let frame = frame::read_response_frame(&mut read_half, &compression)
1346 .await
1347 .map_err(|err| {
1348 warn!("Response reception from {} error: {}", real_addr, err);
1349 WorkerError::NodeDisconnected(real_addr)
1350 })?;
1351
1352 debug!(
1353 "Intercepted Cluster ({}) ({}) -> Driver ({}) frame. opcode: {:?}.",
1354 real_addr,
1355 DisplayableShard(shard),
1356 driver_addr,
1357 &frame.opcode
1358 );
1359
1360 if response_processor_tx.send(frame).is_err() {
1361 warn!("response_processor had exited.");
1362 return Ok::<(), ProxyError>(());
1363 }
1364 }
1365 },
1366 )
1367 .await;
1368 }
1369
1370 async fn sender_to_driver(
1371 self,
1372 mut write_half: impl AsyncWrite + Unpin,
1373 mut responses_rx: mpsc::UnboundedReceiver<ResponseFrame>,
1374 mut connection_close_notifier: ConnectionCloseNotifier,
1375 mut terminate_notifier: TerminateNotifier,
1376 compression: CompressionReader,
1377 ) {
1378 let shard = self.shard;
1379 self.run_until_interrupted(
1380 "sender_to_driver",
1381 |driver_addr, proxy_addr, _real_addr| async move {
1382 loop {
1383 let response = match responses_rx.recv().await {
1384 Some(response) => response,
1385 None => {
1386 if terminate_notifier.try_recv().is_err()
1387 && connection_close_notifier.try_recv().is_err()
1388 {
1389 warn!("Response processor had exited");
1390 }
1391 return Ok(());
1392 }
1393 };
1394
1395 debug!(
1396 "Sending Proxy ({}) ({}) -> Driver ({}) frame. opcode: {:?}.",
1397 proxy_addr,
1398 DisplayableShard(shard),
1399 driver_addr,
1400 &response.opcode
1401 );
1402 if response.write(&mut write_half, &compression).await.is_err() {
1403 if terminate_notifier.try_recv().is_err()
1404 && connection_close_notifier.try_recv().is_err()
1405 {
1406 warn!("Driver dropped connection");
1407 return Err(WorkerError::DriverDisconnected(driver_addr).into());
1408 }
1409 return Ok(());
1410 }
1411 }
1412 },
1413 )
1414 .await;
1415 }
1416
1417 async fn sender_to_cluster(
1418 self,
1419 mut write_half: impl AsyncWrite + Unpin,
1420 mut requests_rx: mpsc::UnboundedReceiver<RequestFrame>,
1421 mut connection_close_notifier: ConnectionCloseNotifier,
1422 mut terminate_notifier: TerminateNotifier,
1423 compression: CompressionReader,
1424 ) {
1425 let shard = self.shard;
1426 self.run_until_interrupted(
1427 "sender_to_driver",
1428 |_driver_addr, proxy_addr, real_addr| async move {
1429 let real_addr = real_addr.expect("BUG: no real_addr in cluster worker");
1430 loop {
1431 let request = match requests_rx.recv().await {
1432 Some(request) => request,
1433 None => {
1434 if terminate_notifier.try_recv().is_err()
1435 && connection_close_notifier.try_recv().is_err()
1436 {
1437 warn!("Request processor had exited");
1438 }
1439 return Ok(());
1440 }
1441 };
1442
1443 debug!(
1444 "Sending Proxy ({}) -> Cluster ({}) ({}) frame. opcode: {:?}.",
1445 proxy_addr,
1446 real_addr,
1447 DisplayableShard(shard),
1448 &request.opcode
1449 );
1450
1451 if request.write(&mut write_half, &compression).await.is_err() {
1452 if terminate_notifier.try_recv().is_err()
1453 && connection_close_notifier.try_recv().is_err()
1454 {
1455 warn!("Node {} dropped connection", real_addr);
1456 return Err(WorkerError::NodeDisconnected(real_addr).into());
1457 }
1458 return Ok(());
1459 }
1460 }
1461 },
1462 )
1463 .await;
1464 }
1465
1466 #[expect(clippy::too_many_arguments)]
1467 async fn request_processor(
1468 self,
1469 mut requests_rx: mpsc::UnboundedReceiver<RequestFrame>,
1470 driver_tx: mpsc::UnboundedSender<ResponseFrame>,
1471 cluster_tx: mpsc::UnboundedSender<RequestFrame>,
1472 connection_no: usize,
1473 request_rules: Arc<Mutex<Vec<RequestRule>>>,
1474 connection_close_signaler: ConnectionCloseSignaler,
1475 event_registered_flag: Arc<AtomicBool>,
1476 compression: CompressionWriter,
1477 cc_event_sender: Arc<Mutex<HashMap<usize, mpsc::UnboundedSender<ResponseFrame>>>>,
1478 ) {
1479 let shard = self.shard;
1480 self.run_until_interrupted("request_processor", |driver_addr, _, real_addr| async move {
1481 'mainloop: loop {
1482 match requests_rx.recv().await {
1483 Some(request) => {
1484 if request.opcode == RequestOpcode::Register {
1485 event_registered_flag.store(true, Ordering::Relaxed);
1486 cc_event_sender.lock().unwrap().insert(connection_no, driver_tx.clone());
1490 info!(
1491 "REGISTER seen on connection {} ({} → {} ({})); registered cc_event_sender",
1492 connection_no,
1493 driver_addr,
1494 DisplayableRealAddrOption(real_addr),
1495 DisplayableShard(shard),
1496 );
1497 } else if request.opcode == RequestOpcode::Startup {
1498 match compression.set_from_startup(&request.body) {
1499 Err(err) => error!("Failed to deserialize STARTUP frame: {}", err),
1500 Ok(read_compression) => info!(
1501 "Intercepted STARTUP frame ({} -> {} ({})), so set compression accordingly to {:?}.",
1502 driver_addr,
1503 DisplayableRealAddrOption(real_addr),
1504 DisplayableShard(shard),
1505 read_compression
1506 )
1507 };
1508 }
1509
1510 let ctx = EvaluationContext {
1511 connection_seq_no: connection_no,
1512 opcode: FrameOpcode::Request(request.opcode),
1513 frame_body: request.body.clone(),
1514 connection_has_events: event_registered_flag.load(Ordering::Relaxed),
1515 };
1516 let mut guard = request_rules.lock().unwrap();
1517 '_ruleloop: for (i, request_rule) in guard.iter_mut().enumerate() {
1518 if request_rule.0.eval(&ctx) {
1519 debug!("Applying rule no={} to request ({} -> {} ({})).", i, driver_addr, DisplayableRealAddrOption(real_addr), DisplayableShard(shard));
1520 debug!("-> Applied rule: {:?}", request_rule);
1521 debug!("-> To request: {:?}", ctx.opcode);
1522 trace!("{:?}", request);
1523
1524 if let Some(ref tx) = request_rule.1.feedback_channel {
1525 tx.send((request.clone(), shard)).unwrap_or_else(|err|
1526 warn!("Could not send received request as feedback: {}", err)
1527 );
1528 }
1529
1530 let request_rule = request_rule.clone();
1531 let to_addressee_action = request_rule.1.to_addressee;
1532 let to_sender_action = request_rule.1.to_sender;
1533 let drop_connection_action = request_rule.1.drop_connection;
1534
1535 let cluster_tx_clone = cluster_tx.clone();
1536 let request_clone = request.clone();
1537 let pass_action = async move {
1538 if let Some(ref pass_action) = to_addressee_action {
1539 if let Some(time) = pass_action.delay {
1540 tokio::time::sleep(time).await;
1541 }
1542 let passed_frame = match pass_action.msg_processor {
1543 Some(ref processor) => processor(request_clone),
1544 None => request_clone,
1545 };
1546 let _ = cluster_tx_clone.send(passed_frame);
1547 };
1548 };
1549
1550 let driver_tx_clone = driver_tx.clone();
1551 let request_clone = request.clone();
1552 let forge_action = async move {
1553 if let Some(ref forge_action) = to_sender_action {
1554 if let Some(time) = forge_action.delay {
1555 tokio::time::sleep(time).await;
1556 }
1557 let forged_frame = {
1558 let processor = forge_action.msg_processor.as_ref()
1559 .expect("Frame processor is required to forge a frame.");
1560 processor(request_clone)
1561 };
1562 let _ = driver_tx_clone.send(forged_frame);
1563 };
1564 };
1565
1566 let connection_close_signaler_clone =
1567 connection_close_signaler.clone();
1568 let drop_action = async move {
1569 if let Some(ref delay) = drop_connection_action {
1570 if let Some(time) = delay {
1571 tokio::time::sleep(*time).await;
1572 }
1573 info!(
1575 "Dropping connection between {} and {} ({}) (as requested by a proxy rule)!",
1576 driver_addr,
1577 DisplayableRealAddrOption(real_addr),
1578 DisplayableShard(shard),
1579 );
1580 let _ = connection_close_signaler_clone.send(());
1581 }
1582 };
1583
1584 tokio::task::spawn(async {
1585 futures::join!(pass_action, forge_action, drop_action);
1586 });
1587
1588 continue 'mainloop; }
1590 }
1591 let _ = cluster_tx.send(request); }
1593 None => {
1594 if event_registered_flag.load(Ordering::Relaxed) {
1598 cc_event_sender.lock().unwrap().remove(&connection_no);
1599 info!(
1600 "Control connection {} ({} → {} ({})) closed; removed cc_event_sender",
1601 connection_no,
1602 driver_addr,
1603 DisplayableRealAddrOption(real_addr),
1604 DisplayableShard(shard),
1605 );
1606 }
1607 return Ok(());
1608 }
1609 }
1610 }
1611 })
1612 .await;
1613 }
1614
1615 #[expect(clippy::too_many_arguments)]
1616 async fn response_processor(
1617 self,
1618 mut responses_rx: mpsc::UnboundedReceiver<ResponseFrame>,
1619 driver_tx: mpsc::UnboundedSender<ResponseFrame>,
1620 cluster_tx: mpsc::UnboundedSender<RequestFrame>,
1621 connection_no: usize,
1622 response_rules: Arc<Mutex<Vec<ResponseRule>>>,
1623 connection_close_signaler: ConnectionCloseSignaler,
1624 event_registered_flag: Arc<AtomicBool>,
1625 ) {
1626 let shard = self.shard;
1627 self.run_until_interrupted("response_processor", |driver_addr, _, real_addr| async move {
1628 'mainloop: loop {
1629 match responses_rx.recv().await {
1630 Some(response) => {
1631 let ctx = EvaluationContext {
1632 connection_seq_no: connection_no,
1633 opcode: FrameOpcode::Response(response.opcode),
1634 frame_body: response.body.clone(),
1635 connection_has_events: event_registered_flag.load(Ordering::Relaxed),
1636 };
1637 let mut guard = response_rules.lock().unwrap();
1638 '_ruleloop: for (i, response_rule) in guard.iter_mut().enumerate() {
1639 if response_rule.0.eval(&ctx) {
1640 debug!("Applying rule no={} to response ({} -> {} ({})).", i, DisplayableRealAddrOption(real_addr), driver_addr, DisplayableShard(shard));
1641 debug!("-> Applied rule: {:?}", response_rule);
1642 debug!("-> To response: {:?}", ctx.opcode);
1643 trace!("{:?}", response);
1644
1645 if let Some(ref tx) = response_rule.1.feedback_channel {
1646 tx.send((response.clone(), shard)).unwrap_or_else(|err| warn!(
1647 "Could not send received response as feedback: {}", err
1648 ));
1649 }
1650
1651 let response_rule = response_rule.clone();
1652 let to_addressee_action = response_rule.1.to_addressee;
1653 let to_sender_action = response_rule.1.to_sender;
1654 let drop_connection_action = response_rule.1.drop_connection;
1655
1656 let response_clone = response.clone();
1657 let driver_tx_clone = driver_tx.clone();
1658 let pass_action = async move {
1659 if let Some(ref pass_action) = to_addressee_action {
1660 if let Some(time) = pass_action.delay {
1661 tokio::time::sleep(time).await;
1662 }
1663 let passed_frame = match pass_action.msg_processor {
1664 Some(ref processor) => processor(response_clone),
1665 None => response_clone,
1666 };
1667 let _ = driver_tx_clone.send(passed_frame);
1668 };
1669 };
1670
1671 let response_clone = response.clone();
1672 let cluster_tx_clone = cluster_tx.clone();
1673 let forge_action = async move {
1674 if let Some(ref forge_action) = to_sender_action {
1675 if let Some(time) = forge_action.delay {
1676 tokio::time::sleep(time).await;
1677 }
1678 let forged_frame = {
1679 let processor = forge_action.msg_processor.as_ref()
1680 .expect("Frame processor is required to forge a frame.");
1681 processor(response_clone)
1682 };
1683 let _ = cluster_tx_clone.send(forged_frame);
1684 };
1685 };
1686
1687 let connection_close_signaler_clone =
1688 connection_close_signaler.clone();
1689 let drop_action = async move {
1690 if let Some(ref delay) = drop_connection_action {
1691 if let Some(time) = delay {
1692 tokio::time::sleep(*time).await;
1693 }
1694 info!(
1696 "Dropping connection between {} and {} ({}) (as requested by a proxy rule)!",
1697 driver_addr,
1698 real_addr.expect("BUG: response rules are unavailable for dry-mode proxy!"),
1699 DisplayableShard(shard)
1700 );
1701 let _ = connection_close_signaler_clone.send(());
1702 }
1703 };
1704
1705 tokio::task::spawn(async {
1706 futures::join!(pass_action, forge_action, drop_action);
1707 });
1708
1709 continue 'mainloop;
1710 }
1711 }
1712 let _ = driver_tx.send(response); }
1714 None => return Ok(()),
1715 }
1716 }
1717 })
1718 .await
1719 }
1720}
1721
1722fn open_shard_aware_socket_to_real_node(
1723 real_addr: SocketAddr,
1724 source_port: u16,
1725) -> std::io::Result<(TcpSocket, SocketAddr)> {
1726 let (socket, unspecified_ip) = match real_addr.ip() {
1727 IpAddr::V4(_) => (
1728 TcpSocket::new_v4()?,
1729 IpAddr::V4(std::net::Ipv4Addr::UNSPECIFIED),
1730 ),
1731 IpAddr::V6(_) => (
1732 TcpSocket::new_v6()?,
1733 IpAddr::V6(std::net::Ipv6Addr::UNSPECIFIED),
1734 ),
1735 };
1736
1737 Ok((socket, SocketAddr::new(unspecified_ip, source_port)))
1738}
1739
1740pub fn get_exclusive_local_address() -> IpAddr {
1743 match std::env::var("NEXTEST_TEST_GLOBAL_SLOT") {
1744 Ok(slot) => {
1745 let slot: u16 = slot
1746 .parse()
1747 .unwrap_or_else(|e| panic!("Invalid slot {e:?}"));
1748 get_exclusive_local_address_nextest(slot)
1749 }
1750 Err(VarError::NotPresent) => get_exclusive_local_address_libtest(),
1751 Err(VarError::NotUnicode(e)) => panic!("Invalid slot {e:?}"),
1752 }
1753}
1754
1755fn get_exclusive_local_address_libtest() -> IpAddr {
1756 static ADDRESS_LOWER_THREE_OCTETS: AtomicU32 = AtomicU32::new(4242);
1758 let next_addr = ADDRESS_LOWER_THREE_OCTETS.fetch_add(1, Ordering::Relaxed);
1759 if next_addr > (u32::MAX >> 8) {
1760 panic!("Loopback address pool for tests depleted");
1761 }
1762 let next_addr_bytes = next_addr.to_le_bytes();
1763 IpAddr::V4(Ipv4Addr::new(
1764 127,
1765 next_addr_bytes[2],
1766 next_addr_bytes[1],
1767 next_addr_bytes[0],
1768 ))
1769}
1770
1771fn get_exclusive_local_address_nextest(slot: u16) -> IpAddr {
1772 static ADDRESS_LOWER_OCTET: AtomicU8 = AtomicU8::new(255);
1773 const FREE_RANGES: u16 = 16;
1776 let next_address_lower = ADDRESS_LOWER_OCTET.fetch_sub(1, Ordering::Relaxed);
1777 if next_address_lower == 0 {
1778 panic!("Loopback address pool for this test depleted");
1779 }
1780
1781 let next_range_bytes: [u8; 2] = slot
1782 .checked_add(FREE_RANGES)
1783 .unwrap_or_else(|| panic!("Loopback address pool for tests depleted"))
1784 .to_le_bytes();
1785
1786 IpAddr::V4(Ipv4Addr::new(
1787 127,
1788 next_range_bytes[1],
1789 next_range_bytes[0],
1790 next_address_lower,
1791 ))
1792}
1793
1794#[cfg(test)]
1795mod tests {
1796 use super::compression::no_compression;
1797 use super::*;
1798 use crate::errors::ReadFrameError;
1799 use crate::frame::{FrameType, read_frame, read_request_frame, read_response_frame};
1800 use crate::proxy::compression::with_compression;
1801 use crate::{
1802 Condition, Reaction as _, RequestReaction, ResponseOpcode, ResponseReaction, setup_tracing,
1803 };
1804 use assert_matches::assert_matches;
1805 use bytes::{BufMut, BytesMut};
1806 use futures::future::{join, join3};
1807 use rand::RngCore;
1808 use scylla_cql::frame::request::options;
1809 use scylla_cql::frame::request::{SerializableRequest as _, Startup};
1810 use scylla_cql::frame::types::write_string_multimap;
1811 use scylla_cql::frame::{Compression, flag};
1812 use std::collections::HashMap;
1813 use std::mem;
1814 use std::str::FromStr;
1815 use std::time::Duration;
1816 use tokio::io::{AsyncReadExt as _, AsyncWriteExt as _};
1817 use tokio::sync::oneshot;
1818
1819 #[test]
1820 fn open_shard_aware_socket_to_real_node_uses_real_node_ip_family() {
1821 for (real_addr, unspecified_ip) in [
1822 (
1823 SocketAddr::from(([127, 0, 0, 1], 9042)),
1824 IpAddr::V4(Ipv4Addr::UNSPECIFIED),
1825 ),
1826 (
1827 SocketAddr::from(([0, 0, 0, 0, 0, 0, 0, 1], 9042)),
1828 IpAddr::V6(std::net::Ipv6Addr::UNSPECIFIED),
1829 ),
1830 ] {
1831 let (socket, bind_addr) = loop {
1832 let port_probe =
1833 std::net::TcpListener::bind(SocketAddr::new(unspecified_ip, 0)).unwrap();
1834 let source_port = port_probe.local_addr().unwrap().port();
1835 drop(port_probe);
1836
1837 let (socket, bind_addr) =
1838 open_shard_aware_socket_to_real_node(real_addr, source_port).unwrap();
1839 assert_eq!(bind_addr, SocketAddr::new(unspecified_ip, source_port));
1840 match socket.bind(bind_addr) {
1841 Ok(()) => break (socket, bind_addr),
1842 Err(error) if error.kind() == std::io::ErrorKind::AddrInUse => continue,
1843 Err(error) => panic!("failed to bind shard-aware socket: {error}"),
1844 }
1845 };
1846 assert_eq!(socket.local_addr().unwrap(), bind_addr);
1847 assert_eq!(socket.local_addr().unwrap().is_ipv4(), real_addr.is_ipv4());
1848 }
1849 }
1850
1851 #[tokio::test]
1852 async fn shard_aware_proxy_connects_to_ipv6_node_from_ipv4_listener() {
1853 setup_tracing();
1854 let mock_node_listener = TcpListener::bind("[::1]:0").await.unwrap();
1855 let real_addr = mock_node_listener.local_addr().unwrap();
1856 let (proxy_addr, running_proxy) = loop {
1857 let proxy_port_probe = std::net::TcpListener::bind("127.0.0.1:0").unwrap();
1858 let proxy_addr = proxy_port_probe.local_addr().unwrap();
1859 drop(proxy_port_probe);
1860 let proxy = Proxy::new([Node::new(
1861 real_addr,
1862 proxy_addr,
1863 ShardAwareness::FixedNum(1),
1864 None,
1865 None,
1866 )]);
1867 match proxy.run().await {
1868 Ok(running_proxy) => break (proxy_addr, running_proxy),
1869 Err(DoorkeeperError::DriverConnectionAttempt(_, error))
1870 if error.kind() == std::io::ErrorKind::AddrInUse => {}
1871 Err(error) => panic!("failed to start mixed-family proxy: {error}"),
1872 }
1873 };
1874
1875 let connect_driver = TcpStream::connect(proxy_addr);
1876 let accept_node = mock_node_listener.accept();
1877 let (driver, node) = tokio::join!(connect_driver, accept_node);
1878 let _driver = driver.unwrap();
1879 let (_node, peer_addr) = node.unwrap();
1880 assert!(peer_addr.is_ipv6());
1881
1882 running_proxy.finish().await.unwrap();
1883 }
1884
1885 fn random_body() -> Bytes {
1886 let body_len = (rand::random::<u32>() % 1000) as usize;
1887 let mut body = BytesMut::zeroed(body_len);
1888 rand::rng().fill_bytes(body.as_mut());
1889 body.freeze()
1890 }
1891
1892 async fn respond_with_supported(
1893 conn: &mut TcpStream,
1894 supported_options: &HashMap<String, Vec<String>>,
1895 compression: &CompressionReader,
1896 ) {
1897 let RequestFrame {
1898 params: recvd_params,
1899 opcode: recvd_opcode,
1900 body: recvd_body,
1901 wire_body_len: _,
1902 } = read_request_frame(conn, compression).await.unwrap();
1903 assert_eq!(recvd_params, HARDCODED_OPTIONS_PARAMS);
1904 assert_eq!(recvd_opcode, RequestOpcode::Options);
1905 assert_eq!(recvd_body, Bytes::new()); let mut body = BytesMut::new();
1908 write_string_multimap(supported_options, &mut body).unwrap();
1909
1910 let body = body.freeze();
1911
1912 write_frame(
1913 HARDCODED_OPTIONS_PARAMS.for_response(),
1914 FrameOpcode::Response(ResponseOpcode::Supported),
1915 &body,
1916 conn,
1917 &no_compression(),
1918 )
1919 .await
1920 .unwrap();
1921 }
1922
1923 fn supported_shards_count(shards_count: u16) -> HashMap<String, Vec<String>> {
1924 let mut sharded_info = HashMap::new();
1925 sharded_info.insert(
1926 String::from("SCYLLA_NR_SHARDS"),
1927 vec![shards_count.to_string()],
1928 );
1929 sharded_info
1930 }
1931
1932 fn supported_shard_number(shard_num: TargetShard) -> HashMap<String, Vec<String>> {
1933 let mut sharded_info = HashMap::new();
1934 sharded_info.insert(String::from("SCYLLA_SHARD"), vec![shard_num.to_string()]);
1935 sharded_info
1936 }
1937
1938 async fn respond_with_shards_count(
1939 conn: &mut TcpStream,
1940 shards_count: u16,
1941 compression: &CompressionReader,
1942 ) {
1943 respond_with_supported(conn, &supported_shards_count(shards_count), compression).await;
1944 }
1945
1946 async fn respond_with_shard_num(
1947 conn: &mut TcpStream,
1948 shard_num: TargetShard,
1949 compression: &CompressionReader,
1950 ) {
1951 respond_with_supported(conn, &supported_shard_number(shard_num), compression).await;
1952 }
1953
1954 fn next_local_address_with_port(port: u16) -> SocketAddr {
1955 SocketAddr::new(get_exclusive_local_address(), port)
1956 }
1957
1958 async fn identity_proxy_does_not_mutate_frames(shard_awareness: ShardAwareness) {
1959 let node1_real_addr = next_local_address_with_port(9876);
1960 let node1_proxy_addr = next_local_address_with_port(9876);
1961 let proxy = Proxy::new([Node::new(
1962 node1_real_addr,
1963 node1_proxy_addr,
1964 shard_awareness,
1965 None,
1966 None,
1967 )]);
1968 let running_proxy = proxy.run().await.unwrap();
1969
1970 let mock_node_listener = TcpListener::bind(node1_real_addr).await.unwrap();
1971
1972 let params = FrameParams {
1973 flags: 0,
1974 version: 0x04,
1975 stream: 0,
1976 };
1977 let opcode = FrameOpcode::Request(RequestOpcode::Options);
1978
1979 let body = random_body();
1980
1981 let send_frame_to_shard = async {
1982 let mut conn = TcpStream::connect(node1_proxy_addr).await.unwrap();
1983
1984 write_frame(params, opcode, &body, &mut conn, &no_compression())
1985 .await
1986 .unwrap();
1987 conn
1988 };
1989
1990 let mock_node_action = async {
1991 if let ShardAwareness::QueryNode = shard_awareness {
1992 respond_with_shards_count(
1993 &mut mock_node_listener.accept().await.unwrap().0,
1994 1,
1995 &no_compression(),
1996 )
1997 .await;
1998 }
1999 let (mut conn, _) = mock_node_listener.accept().await.unwrap();
2000 if shard_awareness.is_aware() {
2001 respond_with_shard_num(&mut conn, 1, &no_compression()).await;
2002 }
2003 let RequestFrame {
2004 params: recvd_params,
2005 opcode: recvd_opcode,
2006 body: recvd_body,
2007 wire_body_len: _,
2008 } = read_request_frame(&mut conn, &no_compression())
2009 .await
2010 .unwrap();
2011 assert_eq!(recvd_params, params);
2012 assert_eq!(FrameOpcode::Request(recvd_opcode), opcode);
2013 assert_eq!(recvd_body, body);
2014 conn
2015 };
2016
2017 let (_node_conn, _driver_conn) = join(mock_node_action, send_frame_to_shard).await;
2019 running_proxy.finish().await.unwrap();
2020 }
2021
2022 #[tokio::test]
2023 async fn identity_shard_unaware_proxy_does_not_mutate_frames() {
2024 setup_tracing();
2025 identity_proxy_does_not_mutate_frames(ShardAwareness::Unaware).await
2026 }
2027
2028 #[tokio::test]
2029 async fn identity_shard_aware_proxy_does_not_mutate_frames() {
2030 setup_tracing();
2031 identity_proxy_does_not_mutate_frames(ShardAwareness::QueryNode).await
2032 }
2033
2034 #[tokio::test]
2035 async fn shard_aware_proxy_is_transparent_for_connection_to_shards() {
2036 setup_tracing();
2037 async fn test_for_shards_num(shards_num: u16) {
2038 let node1_real_addr = next_local_address_with_port(9876);
2039 let node1_proxy_addr = next_local_address_with_port(9876);
2040 let proxy = Proxy::new([Node::new(
2041 node1_real_addr,
2042 node1_proxy_addr,
2043 ShardAwareness::FixedNum(shards_num),
2044 None,
2045 None,
2046 )]);
2047 let running_proxy = proxy.run().await.unwrap();
2048
2049 let mock_node_listener = TcpListener::bind(node1_real_addr).await.unwrap();
2050
2051 let (driver_addr_tx, driver_addr_rx) = oneshot::channel::<SocketAddr>();
2052
2053 let send_frame_to_shard = async {
2054 let socket = TcpSocket::new_v4().unwrap();
2055 socket
2056 .bind(SocketAddr::from_str("0.0.0.0:0").unwrap())
2057 .unwrap();
2058 let conn = socket.connect(node1_proxy_addr).await.unwrap();
2059 driver_addr_tx.send(conn.local_addr().unwrap()).unwrap();
2060 conn
2061 };
2062
2063 let mock_node_action = async {
2064 let (conn, remote_addr) = mock_node_listener.accept().await.unwrap();
2065 let driver_addr = driver_addr_rx.await.unwrap();
2066 assert_eq!(
2067 driver_addr.port() % shards_num,
2068 remote_addr.port() % shards_num
2069 );
2070 conn
2071 };
2072
2073 let (_node_conn, _driver_conn) = join(mock_node_action, send_frame_to_shard).await;
2075 running_proxy.finish().await.unwrap();
2076 }
2077
2078 for shard_num in 1..6 {
2079 test_for_shards_num(shard_num).await;
2080 }
2081 }
2082
2083 #[tokio::test]
2084 async fn shard_aware_proxy_queries_shards_number() {
2085 setup_tracing();
2086 async fn test_for_shards_num(shards_num: u16) {
2087 for shard_num in 0..shards_num {
2088 let node1_real_addr = next_local_address_with_port(9876);
2089 let node1_proxy_addr = next_local_address_with_port(9876);
2090 let proxy = Proxy::new([Node::new(
2091 node1_real_addr,
2092 node1_proxy_addr,
2093 ShardAwareness::QueryNode,
2094 None,
2095 None,
2096 )]);
2097 let running_proxy = proxy.run().await.unwrap();
2098
2099 let mock_node_listener = TcpListener::bind(node1_real_addr).await.unwrap();
2100
2101 let (driver_addr_tx, driver_addr_rx) = oneshot::channel::<SocketAddr>();
2102
2103 let mock_driver_addr = next_local_address_with_port(shards_num * 1234 + shard_num);
2104 let send_frame_to_shard = async {
2105 let socket = TcpSocket::new_v4().unwrap();
2106 socket
2107 .bind(mock_driver_addr)
2108 .unwrap_or_else(|_| panic!("driver_addr failed: {mock_driver_addr}"));
2109 driver_addr_tx.send(socket.local_addr().unwrap()).unwrap();
2110 socket.connect(node1_proxy_addr).await.unwrap()
2111 };
2112
2113 let mock_node_action = async {
2114 respond_with_shards_count(
2115 &mut mock_node_listener.accept().await.unwrap().0,
2116 shards_num,
2117 &no_compression(),
2118 )
2119 .await;
2120 let (conn, remote_addr) = mock_node_listener.accept().await.unwrap();
2121 let driver_addr = driver_addr_rx.await.unwrap();
2122 assert_eq!(
2123 driver_addr.port() % shards_num,
2124 remote_addr.port() % shards_num
2125 );
2126 conn
2127 };
2128
2129 let (_node_conn, _driver_conn) = join(mock_node_action, send_frame_to_shard).await;
2130 running_proxy.finish().await.unwrap();
2131 }
2132 }
2133
2134 for shard_num in 1..6 {
2135 test_for_shards_num(shard_num).await;
2136 }
2137 }
2138
2139 #[tokio::test]
2140 async fn forger_proxy_forges_response() {
2141 setup_tracing();
2142 let node1_real_addr = next_local_address_with_port(9876);
2143 let node1_proxy_addr = next_local_address_with_port(9876);
2144
2145 let this_shall_pass = b"This.Shall.Pass.";
2146 let test_msg = b"Test";
2147
2148 let proxy = Proxy::new([Node::new(
2149 node1_real_addr,
2150 node1_proxy_addr,
2151 ShardAwareness::Unaware,
2152 Some(vec![
2153 RequestRule(
2154 Condition::RequestOpcode(RequestOpcode::Register),
2155 RequestReaction::forge_response(Arc::new(|RequestFrame { params, .. }| {
2156 ResponseFrame::new(
2157 params.for_response(),
2158 ResponseOpcode::Event,
2159 Bytes::from_static(test_msg),
2160 )
2161 })),
2162 ),
2163 RequestRule(
2164 Condition::BodyContainsCaseSensitive(Box::new(*this_shall_pass)),
2165 RequestReaction::noop(),
2166 ),
2167 RequestRule(
2168 Condition::True, RequestReaction::forge_response(Arc::new(|RequestFrame { params, .. }| {
2170 ResponseFrame::new(
2171 params.for_response(),
2172 ResponseOpcode::Ready,
2173 Bytes::new(),
2174 )
2175 })),
2176 ),
2177 ]),
2178 None,
2179 )]);
2180 let running_proxy = proxy.run().await.unwrap();
2181
2182 let mock_node_listener = TcpListener::bind(node1_real_addr).await.unwrap();
2183
2184 let params1 = FrameParams {
2185 flags: 2,
2186 version: 0x42,
2187 stream: 42,
2188 };
2189 let opcode1 = FrameOpcode::Request(RequestOpcode::Startup);
2190
2191 let params2 = FrameParams {
2192 flags: 4,
2193 version: 0x04,
2194 stream: 17,
2195 };
2196 let opcode2 = FrameOpcode::Request(RequestOpcode::Register);
2197
2198 let params3 = FrameParams {
2199 flags: 8,
2200 version: 0x04,
2201 stream: 11,
2202 };
2203 let opcode3 = FrameOpcode::Request(RequestOpcode::Execute);
2204
2205 let body1 = random_body();
2206 let body2 = random_body();
2207 let body3 = {
2208 let mut body = BytesMut::new();
2209 body.put(&b"uSeLeSs JuNk"[..]);
2210 body.put(&this_shall_pass[..]);
2211 body.freeze()
2212 };
2213
2214 let send_frame_to_shard = async {
2215 let mut conn = TcpStream::connect(node1_proxy_addr).await.unwrap();
2216
2217 write_frame(params1, opcode1, &body1, &mut conn, &no_compression())
2218 .await
2219 .unwrap();
2220 write_frame(params2, opcode2, &body2, &mut conn, &no_compression())
2221 .await
2222 .unwrap();
2223 write_frame(params3, opcode3, &body3, &mut conn, &no_compression())
2224 .await
2225 .unwrap();
2226
2227 let ResponseFrame {
2228 params: recvd_params,
2229 opcode: recvd_opcode,
2230 body: recvd_body,
2231 wire_body_len: _,
2232 } = read_response_frame(&mut conn, &no_compression())
2233 .await
2234 .unwrap();
2235 assert_eq!(recvd_params, params1.for_response());
2236 assert_eq!(recvd_opcode, ResponseOpcode::Ready);
2237 assert_eq!(recvd_body, Bytes::new());
2238
2239 let ResponseFrame {
2240 params: recvd_params,
2241 opcode: recvd_opcode,
2242 body: recvd_body,
2243 wire_body_len: _,
2244 } = read_response_frame(&mut conn, &no_compression())
2245 .await
2246 .unwrap();
2247 assert_eq!(recvd_params, params2.for_response());
2248 assert_eq!(recvd_opcode, ResponseOpcode::Event);
2249 assert_eq!(recvd_body, Bytes::from_static(test_msg));
2250
2251 conn
2252 };
2253
2254 let mock_node_action = async {
2255 let (mut conn, _) = mock_node_listener.accept().await.unwrap();
2256 let RequestFrame {
2257 params: recvd_params,
2258 opcode: recvd_opcode,
2259 body: recvd_body,
2260 wire_body_len: _,
2261 } = read_request_frame(&mut conn, &no_compression())
2262 .await
2263 .unwrap();
2264 assert_eq!(recvd_params, params3);
2265 assert_eq!(FrameOpcode::Request(recvd_opcode), opcode3);
2266 assert_eq!(recvd_body, body3);
2267
2268 conn
2269 };
2270
2271 let (mut node_conn, mut driver_conn) = join(mock_node_action, send_frame_to_shard).await;
2272
2273 running_proxy.finish().await.unwrap();
2274
2275 assert_matches!(driver_conn.read(&mut [0u8; 1]).await, Ok(0));
2276 assert_matches!(node_conn.read(&mut [0u8; 1]).await, Ok(0));
2277 }
2278
2279 #[tokio::test]
2280 async fn ad_hoc_rules_changing() {
2281 setup_tracing();
2282 let node1_real_addr = next_local_address_with_port(9876);
2283 let node1_proxy_addr = next_local_address_with_port(9876);
2284 let proxy = Proxy::new([Node::new(
2285 node1_real_addr,
2286 node1_proxy_addr,
2287 ShardAwareness::Unaware,
2288 None,
2289 None,
2290 )]);
2291 let mut running_proxy = proxy.run().await.unwrap();
2292
2293 let mock_node_listener = TcpListener::bind(node1_real_addr).await.unwrap();
2294
2295 let params = FrameParams {
2296 flags: 0,
2297 version: 0x04,
2298 stream: 0,
2299 };
2300 let opcode = FrameOpcode::Request(RequestOpcode::Options);
2301
2302 let body = random_body();
2303
2304 let (mut driver, mut node) = {
2305 let results = join(
2306 TcpStream::connect(node1_proxy_addr),
2307 mock_node_listener.accept(),
2308 )
2309 .await;
2310 (results.0.unwrap(), results.1.unwrap().0)
2311 };
2312
2313 async fn request(
2314 driver: &mut TcpStream,
2315 node: &mut TcpStream,
2316 params: FrameParams,
2317 opcode: FrameOpcode,
2318 body: &Bytes,
2319 ) -> Result<RequestFrame, ReadFrameError> {
2320 let (send_res, recv_res) = join(
2321 write_frame(params, opcode, &body.clone(), driver, &no_compression()),
2322 read_request_frame(node, &no_compression()),
2323 )
2324 .await;
2325 send_res.unwrap();
2326 recv_res
2327 }
2328 {
2329 let RequestFrame {
2331 params: recvd_params,
2332 opcode: recvd_opcode,
2333 body: recvd_body,
2334 wire_body_len: _,
2335 } = request(&mut driver, &mut node, params, opcode, &body)
2336 .await
2337 .unwrap();
2338 assert_eq!(recvd_params, params);
2339 assert_eq!(FrameOpcode::Request(recvd_opcode), opcode);
2340 assert_eq!(recvd_body, body);
2341 }
2342 running_proxy.running_nodes[0].change_request_rules(Some(vec![RequestRule(
2343 Condition::True,
2344 RequestReaction::drop_frame(),
2345 )]));
2346
2347 {
2348 tokio::select! {
2350 res = request(&mut driver, &mut node, params, opcode, &body) => panic!("Rules did not work: received response {:?}", res),
2351 _ = tokio::time::sleep(std::time::Duration::from_millis(20)) => (),
2352 };
2353 }
2354
2355 running_proxy.turn_off_rules();
2356
2357 {
2358 let RequestFrame {
2360 params: recvd_params,
2361 opcode: recvd_opcode,
2362 body: recvd_body,
2363 wire_body_len: _,
2364 } = request(&mut driver, &mut node, params, opcode, &body)
2365 .await
2366 .unwrap();
2367 assert_eq!(recvd_params, params);
2368 assert_eq!(FrameOpcode::Request(recvd_opcode), opcode);
2369 assert_eq!(recvd_body, body);
2370 }
2371
2372 running_proxy.finish().await.unwrap();
2373 }
2374
2375 #[tokio::test]
2376 async fn limited_times_condition_expires() {
2377 setup_tracing();
2378 const FAILING_TRIES: usize = 4;
2379 const PASSING_TRIES: usize = 5;
2380
2381 let node1_real_addr = next_local_address_with_port(9876);
2382 let node1_proxy_addr = next_local_address_with_port(9876);
2383 let proxy = Proxy::new([Node::new(
2384 node1_real_addr,
2385 node1_proxy_addr,
2386 ShardAwareness::Unaware,
2387 Some(vec![
2388 RequestRule(
2389 Condition::not(Condition::TrueForLimitedTimes(
2391 FAILING_TRIES + PASSING_TRIES,
2392 )),
2393 RequestReaction::drop_frame(),
2394 ),
2395 RequestRule(
2396 Condition::not(Condition::TrueForLimitedTimes(FAILING_TRIES)),
2398 RequestReaction::noop(),
2399 ),
2400 RequestRule(
2401 Condition::True,
2403 RequestReaction::drop_frame(),
2404 ),
2405 ]),
2406 None,
2407 )]);
2408 let running_proxy = proxy.run().await.unwrap();
2409
2410 let mock_node_listener = TcpListener::bind(node1_real_addr).await.unwrap();
2411
2412 let params = FrameParams {
2413 flags: 0,
2414 version: 0x04,
2415 stream: 0,
2416 };
2417 let opcode = FrameOpcode::Request(RequestOpcode::Options);
2418 let body = random_body();
2419
2420 let (mut driver, mut node) = {
2421 let results = join(
2422 TcpStream::connect(node1_proxy_addr),
2423 mock_node_listener.accept(),
2424 )
2425 .await;
2426 (results.0.unwrap(), results.1.unwrap().0)
2427 };
2428
2429 async fn request(
2430 driver: &mut TcpStream,
2431 node: &mut TcpStream,
2432 params: FrameParams,
2433 opcode: FrameOpcode,
2434 body: &Bytes,
2435 ) -> Result<RequestFrame, ReadFrameError> {
2436 let (send_res, recv_res) = join(
2437 write_frame(params, opcode, &body.clone(), driver, &no_compression()),
2438 read_request_frame(node, &no_compression()),
2439 )
2440 .await;
2441 send_res.unwrap();
2442 recv_res
2443 }
2444
2445 for _ in 0..FAILING_TRIES {
2446 tokio::select! {
2447 res = request(&mut driver, &mut node, params, opcode, &body) => panic!("Rules did not work: received response {:?}", res),
2448 _ = tokio::time::sleep(std::time::Duration::from_millis(10)) => (),
2449 };
2450 }
2451
2452 for _ in 0..PASSING_TRIES {
2453 let RequestFrame {
2454 params: recvd_params,
2455 opcode: recvd_opcode,
2456 body: recvd_body,
2457 wire_body_len: _,
2458 } = request(&mut driver, &mut node, params, opcode, &body)
2459 .await
2460 .unwrap();
2461 assert_eq!(recvd_params, params);
2462 assert_eq!(FrameOpcode::Request(recvd_opcode), opcode);
2463 assert_eq!(recvd_body, body);
2464 }
2465
2466 for _ in 0..3 {
2467 tokio::select! {
2469 res = request(&mut driver, &mut node, params, opcode, &body) => panic!("Rules did not work: received response {:?}", res),
2470 _ = tokio::time::sleep(std::time::Duration::from_millis(10)) => (),
2471 };
2472 }
2473
2474 running_proxy.finish().await.unwrap();
2475 }
2476
2477 #[tokio::test]
2478 async fn proxy_reports_requests_and_responses_as_feedback() {
2479 setup_tracing();
2480 let node1_real_addr = next_local_address_with_port(9876);
2481 let node1_proxy_addr = next_local_address_with_port(9876);
2482
2483 let (request_feedback_tx, mut request_feedback_rx) = mpsc::unbounded_channel();
2484 let (response_feedback_tx, mut response_feedback_rx) = mpsc::unbounded_channel();
2485 let proxy = Proxy::new([Node::new(
2486 node1_real_addr,
2487 node1_proxy_addr,
2488 ShardAwareness::Unaware,
2489 Some(vec![RequestRule(
2490 Condition::True,
2491 RequestReaction::drop_frame().with_feedback_when_performed(request_feedback_tx),
2492 )]),
2493 Some(vec![ResponseRule(
2494 Condition::True,
2495 ResponseReaction::drop_frame().with_feedback_when_performed(response_feedback_tx),
2496 )]),
2497 )]);
2498 let running_proxy = proxy.run().await.unwrap();
2499
2500 let mock_node_listener = TcpListener::bind(node1_real_addr).await.unwrap();
2501
2502 let params = FrameParams {
2503 flags: 0,
2504 version: 0x04,
2505 stream: 0,
2506 };
2507 let request_opcode = FrameOpcode::Request(RequestOpcode::Options);
2508 let response_opcode = FrameOpcode::Response(ResponseOpcode::Ready);
2509
2510 let body = random_body();
2511
2512 let send_frame_to_shard = async {
2513 let mut conn = TcpStream::connect(node1_proxy_addr).await.unwrap();
2514 write_frame(params, request_opcode, &body, &mut conn, &no_compression())
2515 .await
2516 .unwrap();
2517 conn
2518 };
2519
2520 let mock_node_action = async {
2521 let (mut conn, _) = mock_node_listener.accept().await.unwrap();
2522 write_frame(
2523 params.for_response(),
2524 response_opcode,
2525 &body,
2526 &mut conn,
2527 &no_compression(),
2528 )
2529 .await
2530 .unwrap();
2531 conn
2532 };
2533
2534 let (_node_conn, _driver_conn) = join(mock_node_action, send_frame_to_shard).await;
2536
2537 let (feedback_request, _shard) = request_feedback_rx.recv().await.unwrap();
2538 assert_eq!(feedback_request.params, params);
2539 assert_eq!(
2540 FrameOpcode::Request(feedback_request.opcode),
2541 request_opcode
2542 );
2543 assert_eq!(feedback_request.body, body);
2544 let (feedback_response, _shard) = response_feedback_rx.recv().await.unwrap();
2545 assert_eq!(feedback_response.params, params.for_response());
2546 assert_eq!(
2547 FrameOpcode::Response(feedback_response.opcode),
2548 response_opcode
2549 );
2550 assert_eq!(feedback_response.body, body);
2551
2552 running_proxy.finish().await.unwrap();
2553 }
2554
2555 #[tokio::test]
2556 async fn sanity_check_reports_errors() {
2557 setup_tracing();
2558 let node1_real_addr = next_local_address_with_port(9876);
2559 let node1_proxy_addr = next_local_address_with_port(9876);
2560 let proxy = Proxy::new([Node::new(
2561 node1_real_addr,
2562 node1_proxy_addr,
2563 ShardAwareness::Unaware,
2564 None,
2565 None,
2566 )]);
2567 let mut running_proxy = proxy.run().await.unwrap();
2568
2569 let mock_node_listener = TcpListener::bind(node1_real_addr).await.unwrap();
2570
2571 let send_frame_to_shard = async {
2572 let mut conn = TcpStream::connect(node1_proxy_addr).await.unwrap();
2573
2574 conn.write_all(b"uselessJunk").await.unwrap();
2575 conn
2576 };
2577
2578 let mock_node_action = async {
2579 let (conn, _) = mock_node_listener.accept().await.unwrap();
2580 conn
2581 };
2582
2583 let (node_conn, driver_conn) = join(mock_node_action, send_frame_to_shard).await;
2584
2585 running_proxy.sanity_check().unwrap();
2586
2587 mem::drop(driver_conn);
2588 assert_matches!(
2589 running_proxy.wait_for_error().await,
2590 Some(ProxyError::Worker(WorkerError::DriverDisconnected(_)))
2591 );
2592 running_proxy.sanity_check().unwrap();
2593
2594 mem::drop(node_conn);
2595 assert_matches!(
2596 running_proxy.wait_for_error().await,
2597 Some(ProxyError::Worker(WorkerError::NodeDisconnected(_)))
2598 );
2599 running_proxy.sanity_check().unwrap();
2600
2601 let _ = running_proxy.finish().await;
2603 }
2604
2605 #[tokio::test]
2606 async fn proxy_processes_requests_concurrently() {
2607 setup_tracing();
2608 let node1_real_addr = next_local_address_with_port(9876);
2609 let node1_proxy_addr = next_local_address_with_port(9876);
2610
2611 let delay = Duration::from_millis(60);
2612
2613 let proxy = Proxy::new([Node::new(
2614 node1_real_addr,
2615 node1_proxy_addr,
2616 ShardAwareness::Unaware,
2617 Some(vec![RequestRule(
2618 Condition::TrueForLimitedTimes(1),
2619 RequestReaction::delay(delay),
2620 )]),
2621 None,
2622 )]);
2623 let running_proxy = proxy.run().await.unwrap();
2624
2625 let mock_node_listener = TcpListener::bind(node1_real_addr).await.unwrap();
2626
2627 let params1 = FrameParams {
2628 flags: 0,
2629 version: 0x04,
2630 stream: 0,
2631 };
2632 let opcode1 = FrameOpcode::Request(RequestOpcode::Options);
2633
2634 let body1 = random_body();
2635
2636 let params2 = FrameParams {
2637 flags: 0,
2638 version: 0x04,
2639 stream: 0,
2640 };
2641 let opcode2 = FrameOpcode::Request(RequestOpcode::Register);
2642
2643 let body2 = random_body();
2644
2645 let send_frame_to_shard = async {
2646 let mut conn = TcpStream::connect(node1_proxy_addr).await.unwrap();
2647
2648 write_frame(params1, opcode1, &body1, &mut conn, &no_compression())
2649 .await
2650 .unwrap();
2651 write_frame(params2, opcode2, &body2, &mut conn, &no_compression())
2652 .await
2653 .unwrap();
2654 conn
2655 };
2656
2657 let mock_node_action = async {
2658 let (mut conn, _) = mock_node_listener.accept().await.unwrap();
2659 let RequestFrame {
2660 params: recvd_params,
2661 opcode: recvd_opcode,
2662 body: recvd_body,
2663 wire_body_len: _,
2664 } = read_request_frame(&mut conn, &no_compression())
2665 .await
2666 .unwrap();
2667 assert_eq!(recvd_params, params2);
2668 assert_eq!(FrameOpcode::Request(recvd_opcode), opcode2);
2669 assert_eq!(recvd_body, body2);
2670 conn
2671 };
2672
2673 let (_node_conn, _driver_conn) =
2675 tokio::time::timeout(delay, join(mock_node_action, send_frame_to_shard))
2676 .await
2677 .expect("Request processing was not concurrent");
2678 running_proxy.finish().await.unwrap();
2679 }
2680
2681 #[tokio::test]
2682 async fn dry_mode_proxy_drops_incoming_frames() {
2683 setup_tracing();
2684 let node1_proxy_addr = next_local_address_with_port(9876);
2685 let proxy = Proxy::new([Node::new_dry_mode(node1_proxy_addr, None)]);
2686 let running_proxy = proxy.run().await.unwrap();
2687
2688 let params = FrameParams {
2689 flags: 0,
2690 version: 0x04,
2691 stream: 0,
2692 };
2693 let opcode = FrameOpcode::Request(RequestOpcode::Options);
2694
2695 let body = random_body();
2696
2697 let mut conn = TcpStream::connect(node1_proxy_addr).await.unwrap();
2698
2699 write_frame(params, opcode, &body, &mut conn, &no_compression())
2700 .await
2701 .unwrap();
2702 tokio::time::sleep(Duration::from_millis(3)).await;
2704 running_proxy.finish().await.unwrap();
2705 }
2706
2707 #[tokio::test]
2708 async fn dry_mode_forger_proxy_forges_response() {
2709 setup_tracing();
2710 let node1_proxy_addr = next_local_address_with_port(9876);
2711
2712 let this_shall_pass = b"This.Shall.Pass.";
2713 let test_msg = b"Test";
2714
2715 let proxy = Proxy::new([Node::new_dry_mode(
2716 node1_proxy_addr,
2717 Some(vec![
2718 RequestRule(
2719 Condition::RequestOpcode(RequestOpcode::Register),
2720 RequestReaction::forge_response(Arc::new(|RequestFrame { params, .. }| {
2721 ResponseFrame::new(
2722 params.for_response(),
2723 ResponseOpcode::Event,
2724 Bytes::from_static(test_msg),
2725 )
2726 })),
2727 ),
2728 RequestRule(
2729 Condition::BodyContainsCaseSensitive(Box::new(*this_shall_pass)),
2730 RequestReaction::noop(),
2731 ),
2732 RequestRule(
2733 Condition::True, RequestReaction::forge_response(Arc::new(|RequestFrame { params, .. }| {
2735 ResponseFrame::new(
2736 params.for_response(),
2737 ResponseOpcode::Ready,
2738 Bytes::new(),
2739 )
2740 })),
2741 ),
2742 ]),
2743 )]);
2744 let running_proxy = proxy.run().await.unwrap();
2745
2746 let params1 = FrameParams {
2747 flags: 2,
2748 version: 0x42,
2749 stream: 42,
2750 };
2751 let opcode1 = FrameOpcode::Request(RequestOpcode::Startup);
2752
2753 let params2 = FrameParams {
2754 flags: 4,
2755 version: 0x04,
2756 stream: 17,
2757 };
2758 let opcode2 = FrameOpcode::Request(RequestOpcode::Register);
2759
2760 let params3 = FrameParams {
2761 flags: 8,
2762 version: 0x04,
2763 stream: 11,
2764 };
2765 let opcode3 = FrameOpcode::Request(RequestOpcode::Execute);
2766
2767 let body1 = random_body();
2768 let body2 = random_body();
2769 let body3 = {
2770 let mut body = BytesMut::new();
2771 body.put(&b"uSeLeSs JuNk"[..]);
2772 body.put(&this_shall_pass[..]);
2773 body.freeze()
2774 };
2775
2776 let mut conn = TcpStream::connect(node1_proxy_addr).await.unwrap();
2777
2778 write_frame(params1, opcode1, &body1, &mut conn, &no_compression())
2779 .await
2780 .unwrap();
2781 write_frame(params2, opcode2, &body2, &mut conn, &no_compression())
2782 .await
2783 .unwrap();
2784 write_frame(params3, opcode3, &body3, &mut conn, &no_compression())
2785 .await
2786 .unwrap();
2787
2788 let ResponseFrame {
2789 params: recvd_params,
2790 opcode: recvd_opcode,
2791 body: recvd_body,
2792 wire_body_len: _,
2793 } = read_response_frame(&mut conn, &no_compression())
2794 .await
2795 .unwrap();
2796 assert_eq!(recvd_params, params1.for_response());
2797 assert_eq!(recvd_opcode, ResponseOpcode::Ready);
2798 assert_eq!(recvd_body, Bytes::new());
2799
2800 let ResponseFrame {
2801 params: recvd_params,
2802 opcode: recvd_opcode,
2803 body: recvd_body,
2804 wire_body_len: _,
2805 } = read_response_frame(&mut conn, &no_compression())
2806 .await
2807 .unwrap();
2808 assert_eq!(recvd_params, params2.for_response());
2809 assert_eq!(recvd_opcode, ResponseOpcode::Event);
2810 assert_eq!(recvd_body, Bytes::from_static(test_msg));
2811
2812 running_proxy.finish().await.unwrap();
2813
2814 assert_matches!(conn.read(&mut [0u8; 1]).await, Ok(0));
2815 }
2816
2817 #[tokio::test]
2821 async fn proxy_reports_target_shard_as_feedback() {
2822 setup_tracing();
2823
2824 let node_port = 10101;
2825 let node_real_addr = next_local_address_with_port(node_port);
2826 let mock_node_listener = TcpListener::bind(node_real_addr).await.unwrap();
2827
2828 let params = FrameParams {
2829 flags: 0,
2830 version: 0x04,
2831 stream: 0,
2832 };
2833 let request_opcode = FrameOpcode::Request(RequestOpcode::Options);
2834 let response_opcode = FrameOpcode::Response(ResponseOpcode::Ready);
2835
2836 let body = random_body();
2837
2838 for shards_count in 2..9 {
2839 let driver1_shard = shards_count - 1;
2841 let driver2_shard = shards_count - 2;
2842 let node_proxy_addr = next_local_address_with_port(node_port);
2843
2844 let (request_feedback_tx, mut request_feedback_rx) = mpsc::unbounded_channel();
2845 let (response_feedback_tx, mut response_feedback_rx) = mpsc::unbounded_channel();
2846
2847 let proxy = Proxy::new([Node::new(
2848 node_real_addr,
2849 node_proxy_addr,
2850 ShardAwareness::FixedNum(shards_count),
2851 Some(vec![RequestRule(
2852 Condition::True,
2853 RequestReaction::drop_frame().with_feedback_when_performed(request_feedback_tx),
2854 )]),
2855 Some(vec![ResponseRule(
2856 Condition::True,
2857 ResponseReaction::drop_frame()
2858 .with_feedback_when_performed(response_feedback_tx),
2859 )]),
2860 )]);
2861 let running_proxy = proxy.run().await.unwrap();
2862
2863 fn draw_source_port_for_shard(shards_count: u16, shard: u16) -> u16 {
2865 assert!(shard < shards_count);
2866 49152u16.next_multiple_of(shards_count) + shard
2867 }
2868
2869 async fn bind_socket_for_shard(shards_count: u16, shard: u16) -> TcpSocket {
2870 let socket = TcpSocket::new_v4().unwrap();
2871 let initial_port = draw_source_port_for_shard(shards_count, shard);
2872
2873 let mut desired_addr =
2874 SocketAddr::new(IpAddr::V4(Ipv4Addr::new(0, 0, 0, 0)), initial_port);
2875 while socket.bind(desired_addr).is_err() {
2876 let next_port = desired_addr.port().wrapping_add(shards_count);
2878 if next_port == initial_port {
2879 panic!("No more ports left");
2880 }
2881 desired_addr.set_port(next_port);
2882 }
2883
2884 socket
2885 }
2886
2887 let body_ref = &body;
2888 let send_frame_to_shard = |driver_shard: u16| async move {
2889 let socket = bind_socket_for_shard(shards_count, driver_shard).await;
2890 let mut conn = socket.connect(node_proxy_addr).await.unwrap();
2891
2892 write_frame(
2893 params,
2894 request_opcode,
2895 body_ref,
2896 &mut conn,
2897 &no_compression(),
2898 )
2899 .await
2900 .unwrap();
2901 conn
2902 };
2903
2904 let mock_driver1_action = send_frame_to_shard(driver1_shard);
2905 let mock_driver2_action = send_frame_to_shard(driver2_shard);
2906
2907 let mock_node_action = async {
2909 let mut conns_futs = (0..2)
2910 .map(|_| async {
2911 let (mut conn, driver_addr) = mock_node_listener.accept().await.unwrap();
2912 respond_with_shard_num(
2913 &mut conn,
2914 driver_addr.port() % shards_count,
2915 &no_compression(),
2916 )
2917 .await;
2918 write_frame(
2919 params.for_response(),
2920 response_opcode,
2921 body_ref,
2922 &mut conn,
2923 &no_compression(),
2924 )
2925 .await
2926 .unwrap();
2927 conn
2928 })
2929 .collect::<Vec<_>>();
2930 let conn2 = conns_futs.pop().unwrap().await;
2931 let conn1 = conns_futs.pop().unwrap().await;
2932 (conn1, conn2)
2933 };
2934
2935 let (_node_conns, _driver1_conn, _driver2_conn) =
2937 join3(mock_node_action, mock_driver1_action, mock_driver2_action).await;
2938
2939 let assert_feedback_request = |feedback_request: RequestFrame| {
2940 assert_eq!(feedback_request.params, params);
2941 assert_eq!(
2942 FrameOpcode::Request(feedback_request.opcode),
2943 request_opcode
2944 );
2945 assert_eq!(feedback_request.body, body);
2946 };
2947
2948 let assert_feedback_response = |feedback_response: ResponseFrame| {
2949 assert_eq!(feedback_response.params, params.for_response());
2950 assert_eq!(
2951 FrameOpcode::Response(feedback_response.opcode),
2952 response_opcode
2953 );
2954 assert_eq!(feedback_response.body, body);
2955 };
2956
2957 let (feedback_request, shard1) = request_feedback_rx.recv().await.unwrap();
2958 assert_feedback_request(feedback_request);
2959 let (feedback_request, shard2) = request_feedback_rx.recv().await.unwrap();
2960 assert_feedback_request(feedback_request);
2961 let (feedback_response, shard3) = response_feedback_rx.recv().await.unwrap();
2962 assert_feedback_response(feedback_response);
2963 let (feedback_response, shard4) = response_feedback_rx.recv().await.unwrap();
2964 assert_feedback_response(feedback_response);
2965
2966 let mut expected_shards = [driver1_shard, driver1_shard, driver2_shard, driver2_shard];
2968 expected_shards.sort_unstable();
2969
2970 let mut got_shards = [
2971 shard1.unwrap(),
2972 shard2.unwrap(),
2973 shard3.unwrap(),
2974 shard4.unwrap(),
2975 ];
2976 got_shards.sort_unstable();
2977
2978 assert_eq!(expected_shards, got_shards);
2979
2980 running_proxy.finish().await.unwrap();
2981 }
2982 }
2983
2984 #[tokio::test]
2985 async fn proxy_ignores_control_connection_messages() {
2986 setup_tracing();
2987 let node1_real_addr = next_local_address_with_port(9876);
2988 let node1_proxy_addr = next_local_address_with_port(9876);
2989
2990 let (request_feedback_tx, mut request_feedback_rx) = mpsc::unbounded_channel();
2991 let (response_feedback_tx, mut response_feedback_rx) = mpsc::unbounded_channel();
2992 let proxy = Proxy::new([Node::new(
2993 node1_real_addr,
2994 node1_proxy_addr,
2995 ShardAwareness::Unaware,
2996 Some(vec![RequestRule(
2997 Condition::not(Condition::ConnectionRegisteredAnyEvent),
2998 RequestReaction::noop().with_feedback_when_performed(request_feedback_tx),
2999 )]),
3000 Some(vec![ResponseRule(
3001 Condition::not(Condition::ConnectionRegisteredAnyEvent),
3002 ResponseReaction::noop().with_feedback_when_performed(response_feedback_tx),
3003 )]),
3004 )]);
3005 let running_proxy = proxy.run().await.unwrap();
3006
3007 let mock_node_listener = TcpListener::bind(node1_real_addr).await.unwrap();
3008
3009 let (mut client_socket, mut server_socket) = join(
3010 async { TcpStream::connect(node1_proxy_addr).await.unwrap() },
3011 async { mock_node_listener.accept().await.unwrap().0 },
3012 )
3013 .await;
3014
3015 async fn perform_reqest_response<'a>(
3016 req_opcode: RequestOpcode,
3017 resp_opcode: ResponseOpcode,
3018 client_socket_ref: &'a mut TcpStream,
3019 server_socket_ref: &'a mut TcpStream,
3020 body_base: &'a str,
3021 ) {
3022 let params = FrameParams {
3023 flags: 0,
3024 version: 0x04,
3025 stream: 0,
3026 };
3027
3028 write_frame(
3029 params,
3030 FrameOpcode::Request(req_opcode),
3031 (body_base.to_string() + "|request|").as_bytes(),
3032 client_socket_ref,
3033 &no_compression(),
3034 )
3035 .await
3036 .unwrap();
3037
3038 let received_request =
3039 read_frame(server_socket_ref, FrameType::Request, &no_compression())
3040 .await
3041 .unwrap();
3042 assert_eq!(received_request.1, FrameOpcode::Request(req_opcode));
3043
3044 write_frame(
3045 params.for_response(),
3046 FrameOpcode::Response(resp_opcode),
3047 (body_base.to_string() + "|response|").as_bytes(),
3048 server_socket_ref,
3049 &no_compression(),
3050 )
3051 .await
3052 .unwrap();
3053
3054 let received_response =
3055 read_frame(client_socket_ref, FrameType::Response, &no_compression())
3056 .await
3057 .unwrap();
3058 assert_eq!(received_response.1, FrameOpcode::Response(resp_opcode));
3059 }
3060
3061 for i in 0..5 {
3063 perform_reqest_response(
3064 RequestOpcode::Query,
3065 ResponseOpcode::Result,
3066 &mut client_socket,
3067 &mut server_socket,
3068 &format!("message_before_{i}"),
3069 )
3070 .await
3071 }
3072
3073 perform_reqest_response(
3074 RequestOpcode::Register,
3075 ResponseOpcode::Result,
3076 &mut client_socket,
3077 &mut server_socket,
3078 "message_register",
3079 )
3080 .await;
3081
3082 for i in 0..5 {
3084 perform_reqest_response(
3085 RequestOpcode::Query,
3086 ResponseOpcode::Result,
3087 &mut client_socket,
3088 &mut server_socket,
3089 &format!("message_after_{i}"),
3090 )
3091 .await
3092 }
3093
3094 running_proxy.finish().await.unwrap();
3095
3096 for _ in 0..5 {
3097 let (feedback_request, _shard) = request_feedback_rx.recv().await.unwrap();
3098 assert_eq!(feedback_request.opcode, RequestOpcode::Query);
3099 let (feedback_response, _shard) = response_feedback_rx.recv().await.unwrap();
3100 assert_eq!(feedback_response.opcode, ResponseOpcode::Result);
3101 }
3102
3103 let _ = request_feedback_rx.try_recv().unwrap_err();
3105 let _ = response_feedback_rx.try_recv().unwrap_err();
3106 }
3107
3108 #[tokio::test]
3109 async fn proxy_compresses_and_decompresses_frames_iff_compression_negotiated() {
3110 setup_tracing();
3111 let node1_real_addr = next_local_address_with_port(9876);
3112 let node1_proxy_addr = next_local_address_with_port(9876);
3113
3114 let (request_feedback_tx, mut request_feedback_rx) = mpsc::unbounded_channel();
3115 let (response_feedback_tx, mut response_feedback_rx) = mpsc::unbounded_channel();
3116 let proxy = Proxy::builder()
3117 .with_node(
3118 Node::builder()
3119 .real_address(node1_real_addr)
3120 .proxy_address(node1_proxy_addr)
3121 .shard_awareness(ShardAwareness::Unaware)
3122 .request_rules(vec![RequestRule(
3123 Condition::True,
3124 RequestReaction::noop().with_feedback_when_performed(request_feedback_tx),
3125 )])
3126 .response_rules(vec![ResponseRule(
3127 Condition::True,
3128 ResponseReaction::noop().with_feedback_when_performed(response_feedback_tx),
3129 )])
3130 .build(),
3131 )
3132 .build();
3133 let running_proxy = proxy.run().await.unwrap();
3134
3135 let mock_node_listener = TcpListener::bind(node1_real_addr).await.unwrap();
3136
3137 const PARAMS_REQUEST_NO_COMPRESSION: FrameParams = FrameParams {
3138 flags: 0,
3139 version: 0x04,
3140 stream: 0,
3141 };
3142 const PARAMS_REQUEST_COMPRESSION: FrameParams = FrameParams {
3143 flags: flag::COMPRESSION,
3144 ..PARAMS_REQUEST_NO_COMPRESSION
3145 };
3146 const PARAMS_RESPONSE_NO_COMPRESSION: FrameParams =
3147 PARAMS_REQUEST_NO_COMPRESSION.for_response();
3148 const PARAMS_RESPONSE_COMPRESSION: FrameParams =
3149 PARAMS_REQUEST_NO_COMPRESSION.for_response();
3150
3151 let make_driver_conn = async { TcpStream::connect(node1_proxy_addr).await.unwrap() };
3152 let make_node_conn = async { mock_node_listener.accept().await.unwrap() };
3153
3154 let (mut driver_conn, (mut node_conn, _)) = join(make_driver_conn, make_node_conn).await;
3155
3156 {
3174 let sent_frame = RequestFrame::new(
3175 PARAMS_REQUEST_NO_COMPRESSION,
3176 RequestOpcode::Query,
3177 random_body(),
3178 );
3179
3180 sent_frame
3181 .write(&mut driver_conn, &no_compression())
3182 .await
3183 .unwrap();
3184
3185 let (captured_frame, _) = request_feedback_rx.recv().await.unwrap();
3186 assert_eq!(captured_frame, sent_frame);
3187
3188 let received_frame = read_request_frame(&mut node_conn, &no_compression())
3189 .await
3190 .unwrap();
3191 assert_eq!(received_frame, sent_frame);
3192 }
3193
3194 {
3197 let sent_frame = ResponseFrame::new(
3198 PARAMS_RESPONSE_NO_COMPRESSION,
3199 ResponseOpcode::Result,
3200 random_body(),
3201 );
3202
3203 sent_frame
3204 .write(&mut node_conn, &no_compression())
3205 .await
3206 .unwrap();
3207
3208 let (captured_frame, _) = response_feedback_rx.recv().await.unwrap();
3209 assert_eq!(captured_frame, sent_frame);
3210
3211 let received_frame = read_response_frame(&mut driver_conn, &no_compression())
3212 .await
3213 .unwrap();
3214 assert_eq!(received_frame, sent_frame);
3215 }
3216
3217 {
3222 let startup_body = Startup {
3223 options: std::iter::once((
3224 options::COMPRESSION.into(),
3225 Compression::Lz4.as_str().into(),
3226 ))
3227 .collect(),
3228 }
3229 .to_bytes()
3230 .unwrap();
3231
3232 let sent_frame = RequestFrame::new(
3233 PARAMS_REQUEST_NO_COMPRESSION,
3234 RequestOpcode::Startup,
3235 startup_body,
3236 );
3237
3238 sent_frame
3239 .write(&mut driver_conn, &no_compression())
3240 .await
3241 .unwrap();
3242
3243 let (captured_frame, _) = request_feedback_rx.recv().await.unwrap();
3244 assert_eq!(captured_frame, sent_frame);
3245
3246 let received_frame = read_request_frame(&mut node_conn, &no_compression())
3247 .await
3248 .unwrap();
3249 assert_eq!(received_frame, sent_frame);
3250 }
3251
3252 {
3255 let sent_frame = RequestFrame::new(
3256 PARAMS_REQUEST_COMPRESSION,
3257 RequestOpcode::Query,
3258 random_body(),
3259 );
3260
3261 sent_frame
3262 .write(&mut driver_conn, &with_compression(Compression::Lz4))
3263 .await
3264 .unwrap();
3265
3266 let (captured_frame, _) = request_feedback_rx.recv().await.unwrap();
3267 assert_eq!(captured_frame, sent_frame);
3268
3269 let received_frame =
3270 read_request_frame(&mut node_conn, &with_compression(Compression::Lz4))
3271 .await
3272 .unwrap();
3273 assert_eq!(received_frame, sent_frame);
3274 }
3275
3276 {
3279 let sent_frame = ResponseFrame::new(
3280 PARAMS_RESPONSE_COMPRESSION,
3281 ResponseOpcode::Result,
3282 random_body(),
3283 );
3284
3285 sent_frame
3286 .write(&mut node_conn, &with_compression(Compression::Lz4))
3287 .await
3288 .unwrap();
3289
3290 let (captured_frame, _) = response_feedback_rx.recv().await.unwrap();
3291 assert_eq!(captured_frame, sent_frame);
3292
3293 let received_frame =
3294 read_response_frame(&mut driver_conn, &with_compression(Compression::Lz4))
3295 .await
3296 .unwrap();
3297 assert_eq!(received_frame, sent_frame);
3298 }
3299
3300 running_proxy.finish().await.unwrap();
3301 }
3302
3303 async fn send_register(driver_conn: &mut TcpStream, node_conn: &mut TcpStream) {
3308 let params = FrameParams {
3309 flags: 0,
3310 version: 0x04,
3311 stream: 0,
3312 };
3313 write_frame(
3314 params,
3315 FrameOpcode::Request(RequestOpcode::Register),
3316 b"",
3317 driver_conn,
3318 &no_compression(),
3319 )
3320 .await
3321 .unwrap();
3322
3323 let _req = read_request_frame(node_conn, &no_compression())
3326 .await
3327 .unwrap();
3328 }
3329
3330 #[tokio::test]
3331 async fn inject_event_to_cc_returns_false_when_no_control_connections() {
3332 setup_tracing();
3333 let node_real_addr = next_local_address_with_port(9876);
3334 let node_proxy_addr = next_local_address_with_port(9876);
3335 let proxy = Proxy::new([Node::new(
3336 node_real_addr,
3337 node_proxy_addr,
3338 ShardAwareness::Unaware,
3339 None,
3340 None,
3341 )]);
3342 let running_proxy = proxy.run().await.unwrap();
3343 let _mock_node_listener = TcpListener::bind(node_real_addr).await.unwrap();
3344
3345 assert!(
3347 !running_proxy.running_nodes[0].inject_event_to_cc(Bytes::from_static(b"test")),
3348 "inject_event_to_cc should return false with no control connections"
3349 );
3350
3351 let _ = running_proxy.finish().await;
3353 }
3354
3355 #[tokio::test]
3356 async fn inject_event_to_cc_returns_false_when_connection_did_not_register() {
3357 setup_tracing();
3358 let node_real_addr = next_local_address_with_port(9876);
3359 let node_proxy_addr = next_local_address_with_port(9876);
3360 let proxy = Proxy::new([Node::new(
3361 node_real_addr,
3362 node_proxy_addr,
3363 ShardAwareness::Unaware,
3364 None,
3365 None,
3366 )]);
3367 let running_proxy = proxy.run().await.unwrap();
3368 let mock_node_listener = TcpListener::bind(node_real_addr).await.unwrap();
3369
3370 let _driver_conn = TcpStream::connect(node_proxy_addr).await.unwrap();
3372 let (_node_conn, _) = mock_node_listener.accept().await.unwrap();
3373
3374 assert!(
3376 !running_proxy.running_nodes[0].inject_event_to_cc(Bytes::from_static(b"test")),
3377 "inject_event_to_cc should return false when no REGISTER was sent"
3378 );
3379
3380 running_proxy.finish().await.unwrap();
3381 }
3382
3383 #[tokio::test]
3384 async fn inject_event_to_cc_delivers_event_after_register() {
3385 setup_tracing();
3386 let node_real_addr = next_local_address_with_port(9876);
3387 let node_proxy_addr = next_local_address_with_port(9876);
3388 let proxy = Proxy::new([Node::new(
3389 node_real_addr,
3390 node_proxy_addr,
3391 ShardAwareness::Unaware,
3392 None,
3393 None,
3394 )]);
3395 let running_proxy = proxy.run().await.unwrap();
3396 let mock_node_listener = TcpListener::bind(node_real_addr).await.unwrap();
3397
3398 let (mut driver_conn, mut node_conn) = join(
3399 async { TcpStream::connect(node_proxy_addr).await.unwrap() },
3400 async { mock_node_listener.accept().await.unwrap().0 },
3401 )
3402 .await;
3403
3404 send_register(&mut driver_conn, &mut node_conn).await;
3406
3407 let event_body = Bytes::from_static(b"injected_event_payload");
3409 assert!(
3410 running_proxy.running_nodes[0].inject_event_to_cc(event_body.clone()),
3411 "inject_event_to_cc should return true after REGISTER"
3412 );
3413
3414 let frame = tokio::time::timeout(
3416 Duration::from_millis(100),
3417 read_response_frame(&mut driver_conn, &no_compression()),
3418 )
3419 .await
3420 .expect("timed out waiting for injected event")
3421 .expect("failed to read injected event frame");
3422
3423 assert_eq!(frame.opcode, ResponseOpcode::Event);
3424 assert_eq!(frame.body, event_body);
3425 assert_eq!(frame.params.stream, -1);
3426
3427 running_proxy.finish().await.unwrap();
3428 }
3429
3430 #[tokio::test]
3431 async fn inject_event_to_cc_prunes_closed_connections() {
3432 setup_tracing();
3433 let node_real_addr = next_local_address_with_port(9876);
3434 let node_proxy_addr = next_local_address_with_port(9876);
3435 let proxy = Proxy::new([Node::new(
3436 node_real_addr,
3437 node_proxy_addr,
3438 ShardAwareness::Unaware,
3439 None,
3440 None,
3441 )]);
3442 let running_proxy = proxy.run().await.unwrap();
3443 let mock_node_listener = TcpListener::bind(node_real_addr).await.unwrap();
3444
3445 let (mut driver_conn1, mut node_conn1) = join(
3447 async { TcpStream::connect(node_proxy_addr).await.unwrap() },
3448 async { mock_node_listener.accept().await.unwrap().0 },
3449 )
3450 .await;
3451 send_register(&mut driver_conn1, &mut node_conn1).await;
3452
3453 let (mut driver_conn2, mut node_conn2) = join(
3455 async { TcpStream::connect(node_proxy_addr).await.unwrap() },
3456 async { mock_node_listener.accept().await.unwrap().0 },
3457 )
3458 .await;
3459 send_register(&mut driver_conn2, &mut node_conn2).await;
3460
3461 assert!(running_proxy.running_nodes[0].inject_event_to_cc(Bytes::from_static(b"ev1")));
3463
3464 let f1 = read_response_frame(&mut driver_conn1, &no_compression())
3466 .await
3467 .unwrap();
3468 let f2 = read_response_frame(&mut driver_conn2, &no_compression())
3469 .await
3470 .unwrap();
3471 assert_eq!(f1.body, Bytes::from_static(b"ev1"));
3472 assert_eq!(f2.body, Bytes::from_static(b"ev1"));
3473
3474 drop(driver_conn1);
3476 drop(node_conn1);
3477
3478 tokio::time::sleep(Duration::from_millis(100)).await;
3481
3482 assert!(running_proxy.running_nodes[0].inject_event_to_cc(Bytes::from_static(b"ev2")));
3485
3486 let f2 = read_response_frame(&mut driver_conn2, &no_compression())
3487 .await
3488 .unwrap();
3489 assert_eq!(f2.body, Bytes::from_static(b"ev2"));
3490
3491 drop(driver_conn2);
3493 drop(node_conn2);
3494
3495 tokio::time::sleep(Duration::from_millis(100)).await;
3496
3497 assert!(
3499 !running_proxy.running_nodes[0].inject_event_to_cc(Bytes::from_static(b"ev3")),
3500 "inject_event_to_cc should return false after all control connections closed"
3501 );
3502
3503 let _ = running_proxy.finish().await;
3506 }
3507}