Skip to main content

scylla_proxy/
proxy.rs

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
26// Used to notify the user that the proxy finished - this happens when all Senders are dropped.
27type FinishWaiter = mpsc::Receiver<()>;
28type FinishGuard = mpsc::Sender<()>;
29
30// Used to tell all the proxy workers to stop when the user requests that with [RunningProxy::finish()].
31type TerminateNotifier = tokio::sync::broadcast::Receiver<()>;
32type TerminateSignaler = tokio::sync::broadcast::Sender<()>;
33
34// Used to tell all proxy workers working on same connection to stop when
35// a rule being applied has connection drop set.
36type ConnectionCloseNotifier = tokio::sync::broadcast::Receiver<()>;
37type ConnectionCloseSignaler = tokio::sync::broadcast::Sender<()>;
38
39// Used to gather errors from all proxy workers and propagate them to the proxy user,
40// returning the first of them from [RunningProxy::finish()].
41type 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/// Specifies proxy's behaviour regarding shard awareness.
51#[derive(Clone, Copy, Debug)]
52pub enum ShardAwareness {
53    /// Acts as if the connection was made to the shard-unaware port.
54    Unaware,
55    /// The first time the driver attempts to connect to the particular node (through proxy),
56    /// the related node is first queried on a temporary connection for its number of shards,
57    /// and only then establishes another connection for the driver's real communication with the node.
58    /// If the queried node does not provide sharding info (e.g. in case of a Cassandra node),
59    /// then this mode behaves as Unaware.
60    QueryNode,
61    /// Binds to the port that is the same as the driver's port modulo the provided number of shards.
62    FixedNum(u16),
63}
64
65impl ShardAwareness {
66    pub fn is_aware(&self) -> bool {
67        !matches!(self, Self::Unaware)
68    }
69}
70
71/// Node can be either Real (truly backed by a Scylla node) or Simulated
72/// (driver believes it's real, but we merely simulate it with the proxy).
73/// In Simulated mode, no node address is provided and proxy does not attempt
74/// to establish connection with a Scylla node.
75///
76/// For Real node, all workers are created, so such frame flow is possible:
77/// [driver] -> receiver_from_driver -> requests_processor -> sender_to_cluster -> [node] (ordinary request flow)
78///                                        |                    /\
79///     (forging response) ++--------------+      +-------------++ (forging request)
80///                        \/                     |
81/// [driver] <- sender_to_driver <- response_processor <- receiver_from_cluster <- [node] (ordinary response flow)
82///
83/// For Simulated node, it looks like this:
84/// [driver] -> receiver_from_driver -> requests_processor -+
85///                                                         |   (forging response)
86/// [driver] <- sender_to_driver <--------------------------+
87///
88/// For Real node, the default reaction to a frame is to pass it to its intended addresse.
89/// For Simulated node, the default reaction to a request is to drop it.
90enum 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    /// Creates an abstract node that is backed by a real Scylla node.
107    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    /// Creates a simulated node that is not backed by any real Scylla node.
126    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    /// Creates an abstract node that is backed by a real Scylla node.
180    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    /// Creates a simulated node that is not backed by any real Scylla node.
193    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    /// Build a translation map based on provided proxy and node addresses.
322    /// The map can be passed to `Session` `address_translator()` to ensure
323    /// that the driver contacts the nodes through the proxy (and not directly).
324    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    /// Runs the [Proxy], i.e. makes it ready for accepting drivers' connections.
342    /// Returns a [RunningProxy] handle that can be used to stop the proxy or change the rules.
343    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?; // await doorkeeper creation, including binding to a socket
388            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
405/// Creates [`DuplexStream`]s that directly connect to the associated proxy's
406/// nodes.
407///
408/// This is an alternative method of establishing a transport layer connection
409/// to the proxy nodes. The purpose of this alternative method is to allow
410/// writing tests that use tokio's paused time feature. The advantage of using
411/// a channel is that, when some data is sent on the channel, the waiting
412/// task will be immediately notified and scheduled for execution. For
413/// system sockets this is not guaranteed and the receiving task may be
414/// notified after a real-time delay, which can trick tokio's timer auto-advance
415/// and make it move the time forward, even if the task on the receiving
416/// side would have been notified earlier in a real scenario.
417pub 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
453/// A handle that can be used to change the rules regarding the particular node.
454pub struct RunningNode {
455    request_rules: Arc<Mutex<Vec<RequestRule>>>,
456    response_rules: Option<Arc<Mutex<Vec<ResponseRule>>>>,
457
458    /// Senders to the driver-facing sockets of all control connections (those
459    /// that sent REGISTER). Keyed by `connection_no` so that each connection
460    /// can be individually removed on close without affecting others.
461    ///
462    /// Populated by `request_processor` when it sees a REGISTER frame,
463    /// entries removed when the corresponding connection closes.
464    ///
465    /// Used by [`inject_event_to_cc`](Self::inject_event_to_cc) to push
466    /// unsolicited EVENT frames to every registered control connection.
467    cc_event_sender: Arc<Mutex<HashMap<usize, mpsc::UnboundedSender<ResponseFrame>>>>,
468}
469
470impl RunningNode {
471    /// Replaces the previous request rules with the new ones.
472    pub fn change_request_rules(&mut self, rules: Option<Vec<RequestRule>>) {
473        *self.request_rules.lock().unwrap() = rules.unwrap_or_default();
474    }
475
476    /// Adds new request rules to the end of the list (so with lowest priority)
477    pub fn append_request_rules(&mut self, mut rules: Vec<RequestRule>) {
478        self.request_rules.lock().unwrap().append(&mut rules);
479    }
480
481    /// Adds new request rules to the beginning of the list (so with highest priority)
482    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    /// Replaces the previous response rules with the new ones.
490    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    /// Adds new response rules to the end of the list (so with lowest priority)
500    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    /// Adds new response rules to the beginning of the list (so with highest priority)
510    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    /// Injects a CQL EVENT frame into all registered control connections.
523    ///
524    /// Builds an EVENT frame (stream = −1, flags = 0, opcode = Event) with the
525    /// supplied `body` and sends it to every driver that has sent REGISTER on
526    /// this node. Dead senders (closed connections) are pruned automatically.
527    ///
528    /// Returns `true` if the frame was successfully enqueued to at least one
529    /// control connection, `false` if no control connections are currently
530    /// registered or all sends failed.
531    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            // Remove dead senders (connection already closed on the
551            // receiving end).
552            ok
553        });
554        any_sent
555    }
556}
557
558/// A handle that can be used to stop the proxy or change the rules.
559pub 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    /// Disables all the rules in the proxy, effectively making it a pass-through-only proxy.
569    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    /// Attempts to fetch the first error that has occurred in proxy since last check.
583    /// If no errors occurred, returns Ok(()).
584    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                // As we haven't awaited finish of all workers yet, there must be a faulty case without proper error handling.
590                Err(ProxyError::SanityCheckFailure)
591            }
592        }
593    }
594
595    /// Waits until an error occurs in proxy. If proxy finishes with no errors occurred, returns Err(()).
596    pub async fn wait_for_error(&mut self) -> Option<ProxyError> {
597        self.error_sink.recv().await
598    }
599
600    /// Returns the [`TransportFactory`] object associated with the proxy.
601    pub fn transport_factory(&self) -> Arc<TransportFactory> {
602        Arc::clone(&self.transport_factory)
603    }
604
605    /// Requests termination of all proxy workers and awaits its completion.
606    /// Returns the first error that occurred in proxy.
607    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        // This to make sure that also workers not-yet-spawned when terminate signal was sent will terminate.
616        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                // As we have already awaited finish of all workers, there must be a logic bug.
628                unreachable!("Worker await logic bug!");
629            }
630        }
631    }
632}
633
634/// A worker corresponding to a particular node. It listens in a loop for driver's connections
635/// on specified proxy bind address, respects ports regarding advanced shard-awareness (if set),
636/// to this end obtaining number of shards from the node (if set), then establishes connection
637/// to the node, spawns workers for this connection and continues to listen.
638struct 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    // The other side of the `duplex_receiver`. Not used directly by the doorkeeper
648    // but kept as a field in order to keep `duplex_receiver` open which allows
649    // to simplify some logic and not care about it getting closed.
650    _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, // temporarily, until Doorkeeper examines its ShardAwareness
691            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                    // If a node offers no sharding info, change proxy ShardAwareness to Unaware.
749                    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                // Doorkeeper (which is passed by reference in `self`) keeps
969                // the send half of the `duplex_receiver` alive as one of its
970                // fields, just so that we can avoid dealing with the doorkeeper
971                // becoming closed here
972                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                    // in search for a port that translates to the desired shard
998                    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        // If ShardAwareness is aware (QueryNode or FixedNum variants) and the
1027        // proxy succeeded to know the shards count (in FixedNum we get it for
1028        // free, in QueryNode the initial Options query succeeded and Supported
1029        // contained SCYLLA_SHARDS_NUM), then upon opening each connection to the
1030        // node, the proxy issues another Options requests and acknowledges the
1031        // shard it got connected to.
1032        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        /// Body Snap compression failed.
1124        #[error("Snap compression error: {0}")]
1125        SnapCompressError(Arc<dyn Error + Sync + Send>),
1126
1127        /// Frame is to be compressed, but no compression was negotiated for the connection.
1128        #[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    /// The write end of compression config for a connection.
1135    ///
1136    /// Used by the request processor upon STARTUP frame captured
1137    /// and compression setting retrieved from it.
1138    #[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    /// The read end of compression config for a connection.
1169    ///
1170    /// Used by frame (de)serializers.
1171    #[derive(Debug, Clone)]
1172    pub(crate) struct CompressionReader(CompressionInfo);
1173    impl CompressionReader {
1174        /// Return the compression negotiated for the connection.
1175        ///
1176        /// Outer Option signifies whether the negotiation took place,
1177        /// inner Option is the compression (or lack of it) negotiated.
1178        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    // Compression explicitly turned off.
1242    pub(crate) fn no_compression() -> CompressionReader {
1243        mock_compression_reader(None)
1244    }
1245
1246    // Compression explicitly turned on.
1247    #[cfg(test)] // Currently only used for tests.
1248    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                    // error_propagator could be a field
1289                    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                            // Expose this connection's driver-facing sender so
1487                            // that RunningNode::inject_event_to_cc() can push
1488                            // unsolicited EVENT frames to this control connection.
1489                            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                                        // close connection.
1574                                        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; // only one rule can be applied to one frame
1589                            }
1590                        }
1591                        let _ = cluster_tx.send(request); // default action
1592                    }
1593                    None => {
1594                        // Connection closed. If this was the control
1595                        // connection (REGISTER was seen), remove only this
1596                        // connection's sender from the shared map.
1597                        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                                        // close connection.
1695                                        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); // default action
1713                    }
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
1740// Returns next free IP address for another proxy instance.
1741// Useful for concurrent testing.
1742pub 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    // A big enough number reduces possibility of clashes with user-taken addresses:
1757    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    // This is a heuristic to avoid using low addresses, which I think have
1774    // a higher chance of being taken.
1775    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()); // body should be empty
1906
1907        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        // we keep the connections open until proxy finishes to let it perform clean exit with no disconnects
2018        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            // we keep the connections open until proxy finishes to let it perform clean exit with no disconnects
2074            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, // only the first matching rule is applied, so "True" covers all remaining cases
2169                    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            // one run still without custom rules
2330            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            // one run with custom rules
2349            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            // one run already without custom rules
2359            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                    // this will be always fired after first PASSING_TRIES + FAILING_TRIES
2390                    Condition::not(Condition::TrueForLimitedTimes(
2391                        FAILING_TRIES + PASSING_TRIES,
2392                    )),
2393                    RequestReaction::drop_frame(),
2394                ),
2395                RequestRule(
2396                    // this will be fired for PASSING_TRIES after first FAILING_TRIES
2397                    Condition::not(Condition::TrueForLimitedTimes(FAILING_TRIES)),
2398                    RequestReaction::noop(),
2399                ),
2400                RequestRule(
2401                    // this will be fired for first FAILING_TRIES
2402                    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            // any further number of requests should fail
2468            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        // we keep the connections open until proxy finishes to let it perform clean exit with no disconnects
2535        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        // we keep the connections open until proxy finishes to let it perform clean exit with no disconnects
2602        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        // we keep the connections open until proxy finishes to let it perform clean exit with no disconnects
2674        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        // We assert that after sufficiently long time, no error happens inside proxy.
2703        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, // only the first matching rule is applied, so "True" covers all remaining cases
2734                    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    // The test asserts that once a (mock) driver connects to the proxy from some port,
2818    // the proxy will connect to a shard corresponding to that port and that the target
2819    // shard number will be sent through the feedback channel.
2820    #[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            // Two driver connections are simulated, each to a different shard.
2840            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            /// Choose a source port `p` such that `shard == shard_of_source_port(p)`.
2864            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                    // in search for a port that translates to the desired shard
2877                    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            // Accepts two connections and sends a response to each of them.
2908            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            // we keep the connections open until proxy finishes to let it perform clean exit with no disconnects
2936            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            // expected: {driver1_shard request, driver1_shard response, driver2_shard request, driver2_shard response}
2967            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        // Messages before REGISTER should be fed back to channels
3062        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        // Messages after REGISTER should be passed through without feedback
3083        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        // Response to REGISTER and further requests / responses should be ignored
3104        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        /* Outline of the test:
3157         * 1. "driver" sends an uncompressed, e.g., QUERY frame, feedback returns its uncompressed body,
3158         *    and "node" receives the uncompressed frame.
3159         * 2. "node" responds with an uncompressed RESULT frame, feedback returns its uncompressed body,
3160         *    and "driver" receives the uncompressed frame.
3161         * 3. "driver" sends an uncompressed STARTUP frame, feedback returns its uncompressed body,
3162         *    and "node" receives the uncompressed frame. This step also triggers `CompressionWriter::set()`
3163         *    in the proxy, so the associated `CompressionReader`s are notified about it (and can use
3164         *    the negotiated compression algorithm to (de)compress the frames sent in steps 4. and 5.).
3165         * 4. "driver" sends a compressed, e.g., QUERY frame, feedback returns its uncompressed body,
3166         *    and "node" receives the compressed frame.
3167         * 5. "node" responds with a compressed RESULT frame, feedback returns its uncompressed body,
3168         *    and "driver" receives the compressed frame.
3169         */
3170
3171        // 1. "driver" sends an uncompressed, e.g., QUERY frame, feedback returns its uncompressed body,
3172        //    and "node" receives the uncompressed frame.
3173        {
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        // 2. "node" responds with an uncompressed RESULT frame, feedback returns its uncompressed body,
3195        //    and "driver" receives the uncompressed frame.
3196        {
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        // 3. "driver" sends an uncompressed STARTUP frame, feedback returns its uncompressed body,
3218        //    and "node" receives the uncompressed frame. This step also triggers `CompressionWriter::set()`
3219        //    in the proxy, so the associated `CompressionReader`s are notified about it (and can use
3220        //    the negotiated compression algorithm to (de)compress the frames sent in steps 4. and 5.).
3221        {
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        // 4. "driver" sends a compressed, e.g., QUERY frame, feedback returns its uncompressed body,
3253        //    and "node" receives the compressed frame.
3254        {
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        // 5. "node" responds with a compressed RESULT frame, feedback returns its uncompressed body,
3277        //    and "driver" receives the compressed frame.
3278        {
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    /// Helper: send a REGISTER frame from the driver side and wait until
3304    /// the mock node receives it. This is enough for the proxy's
3305    /// request_processor to register the cc_event_sender for that
3306    /// connection — no response is needed.
3307    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        // Wait until the mock node receives the REGISTER so we know the
3324        // proxy has already processed it and registered the sender.
3325        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        // No connections at all — inject should return false.
3346        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        // finish() may report errors because no real node accepted connections.
3352        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        // Connect but do NOT send REGISTER.
3371        let _driver_conn = TcpStream::connect(node_proxy_addr).await.unwrap();
3372        let (_node_conn, _) = mock_node_listener.accept().await.unwrap();
3373
3374        // Connection exists but hasn't sent REGISTER — inject should return false.
3375        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        // Complete the REGISTER handshake so the proxy registers the cc sender.
3405        send_register(&mut driver_conn, &mut node_conn).await;
3406
3407        // Inject an event.
3408        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        // Read the injected frame on the driver side.
3415        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        // Establish first control connection.
3446        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        // Establish second control connection.
3454        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        // Both connections registered — inject should succeed.
3462        assert!(running_proxy.running_nodes[0].inject_event_to_cc(Bytes::from_static(b"ev1")));
3463
3464        // Read event from both.
3465        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        // Close connection 1 (both sides).
3475        drop(driver_conn1);
3476        drop(node_conn1);
3477
3478        // Give the proxy a moment to detect the closed connection and clean up
3479        // its cc_event_sender entry.
3480        tokio::time::sleep(Duration::from_millis(100)).await;
3481
3482        // Inject again — should still succeed via connection 2, and prune
3483        // the dead sender for connection 1.
3484        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        // Close connection 2 as well.
3492        drop(driver_conn2);
3493        drop(node_conn2);
3494
3495        tokio::time::sleep(Duration::from_millis(100)).await;
3496
3497        // Now all control connections are gone — inject should return false.
3498        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        // finish() may report DriverDisconnected errors from the intentionally
3504        // dropped connections — that's expected.
3505        let _ = running_proxy.finish().await;
3506    }
3507}