Skip to main content

mz_service/
transport.rs

1// Copyright Materialize, Inc. and contributors. All rights reserved.
2//
3// Use of this software is governed by the Business Source License
4// included in the LICENSE file.
5//
6// As of the Change Date specified in that file, in accordance with
7// the Business Source License, use of this software will be governed
8// by the Apache License, Version 2.0.
9
10//! The Cluster Transport Protocol (CTP).
11//!
12//! CTP is the protocol used to transmit commands from controllers to replicas and responses from
13//! replicas to controllers. It runs on top of a reliable bidirectional connection stream, as
14//! provided by TCP or UDS, and adds message framing as well as heartbeating.
15//!
16//! CTP supports any message type that implements the serde [`Serialize`] and [`Deserialize`]
17//! traits. Messages are encoded using the [`bincode`] format and then sent over the wire with a
18//! length prefix.
19//!
20//! A CTP server only serves a single client at a time. If a new client connects while a connection
21//! is already established, the previous connection is canceled.
22
23mod metrics;
24#[cfg(test)]
25mod tests;
26
27use std::convert::Infallible;
28use std::fmt::Debug;
29use std::time::Duration;
30
31use anyhow::bail;
32use async_trait::async_trait;
33use bincode::Options;
34use futures::future;
35use mz_ore::cast::CastInto;
36use mz_ore::netio::{Listener, SocketAddr, Stream, TimedReader, TimedWriter};
37use mz_ore::task::{AbortOnDropHandle, JoinHandle};
38use semver::Version;
39use serde::de::DeserializeOwned;
40use serde::{Deserialize, Serialize};
41use tokio::io::{AsyncRead, AsyncReadExt, AsyncWrite, AsyncWriteExt};
42use tokio::sync::{mpsc, oneshot, watch};
43use tracing::{Instrument, info, info_span, trace, warn};
44
45use crate::client::{GenericClient, Partitionable, Partitioned};
46
47pub use metrics::{ClusterServerMetrics, Metrics, NoopMetrics, PerClusterServerMetrics};
48
49/// Trait for messages that can be sent over CTP.
50pub trait Message: Debug + Send + Sync + Serialize + DeserializeOwned + 'static {}
51impl<T: Debug + Send + Sync + Serialize + DeserializeOwned + 'static> Message for T {}
52
53/// A client for a CTP connection.
54#[derive(Debug)]
55pub struct Client<Out, In> {
56    conn: Connection<Out, In>,
57}
58
59impl<Out: Message, In: Message> Client<Out, In> {
60    /// Connect to the server at the given address.
61    ///
62    /// This call resolves once a connection with the server host was either established, was
63    /// rejected, or timed out.
64    pub async fn connect(
65        address: &str,
66        version: Version,
67        connect_timeout: Duration,
68        idle_timeout: Duration,
69        metrics: impl Metrics<Out, In>,
70    ) -> anyhow::Result<Self> {
71        let dest_host = host_from_address(address);
72        let stream = mz_ore::future::timeout(connect_timeout, Stream::connect(address)).await?;
73        info!(%address, "ctp: connected to server");
74
75        // The connection tasks inherit this span, so their log lines name the peer. A controller
76        // does not otherwise put its replica in scope, and under Kubernetes the address names the
77        // cluster, replica, generation, and process.
78        let span = info_span!("ctp", %address);
79        let conn = Connection::start(stream, version, dest_host, idle_timeout, metrics)
80            .instrument(span)
81            .await?;
82        Ok(Self { conn })
83    }
84}
85
86/// Helper function to extract the host part from an address string.
87///
88/// This function assumes addresses to be of the form `<host>:<port>` or `<protocol>:<host>:<port>`
89/// and yields `None` otherwise.
90fn host_from_address(address: &str) -> Option<String> {
91    let mut p = address.split(':');
92    let (host, port) = match (p.next(), p.next(), p.next(), p.next()) {
93        (Some(host), Some(port), None, None) => (host, port),
94        (Some(_protocol), Some(host), Some(port), None) => (host, port),
95        _ => return None,
96    };
97
98    let _: u16 = port.parse().ok()?;
99    Some(host.into())
100}
101
102impl<Out, In> Client<Out, In>
103where
104    Out: Message,
105    In: Message,
106    (Out, In): Partitionable<Out, In>,
107{
108    /// Create a `Partitioned` client that connects through CTP.
109    pub async fn connect_partitioned(
110        addresses: Vec<String>,
111        version: Version,
112        connect_timeout: Duration,
113        idle_timeout: Duration,
114        metrics: impl Metrics<Out, In>,
115    ) -> anyhow::Result<Partitioned<Self, Out, In>> {
116        let connects = addresses.iter().map(|addr| {
117            Self::connect(
118                addr,
119                version.clone(),
120                connect_timeout,
121                idle_timeout,
122                metrics.clone(),
123            )
124        });
125        let clients = future::try_join_all(connects).await?;
126        Ok(Partitioned::new(clients))
127    }
128}
129
130#[async_trait]
131impl<Out: Message, In: Message> GenericClient<Out, In> for Client<Out, In> {
132    async fn send(&mut self, cmd: Out) -> anyhow::Result<()> {
133        self.conn.send(cmd).await
134    }
135
136    /// # Cancel safety
137    ///
138    /// This method is cancel safe.
139    async fn recv(&mut self) -> anyhow::Result<Option<In>> {
140        // `Connection::recv` is documented to be cancel safe.
141        self.conn.recv().await.map(Some)
142    }
143}
144
145/// Spawn a CTP server that serves connections at the given address.
146pub async fn serve<In, Out, H>(
147    address: SocketAddr,
148    version: Version,
149    server_fqdn: Option<String>,
150    idle_timeout: Duration,
151    handler_fn: impl Fn() -> H,
152    metrics: impl Metrics<Out, In>,
153) -> anyhow::Result<()>
154where
155    In: Message,
156    Out: Message,
157    H: GenericClient<In, Out> + 'static,
158{
159    // Keep a handle to the task serving the current connection, as well as a cancelation token, so
160    // we can cancel it when a new client connects.
161    //
162    // Note that we cannot simply abort the previous connection task because its future isn't known
163    // to be cancel safe. Instead we pass the connection tasks a cancelation token and wait for
164    // them to shut themselves down gracefully once the token gets dropped.
165    let mut connection_task: Option<(JoinHandle<()>, oneshot::Sender<()>)> = None;
166
167    let listener = Listener::bind(&address).await?;
168    info!(%address, "ctp: listening for client connections");
169
170    loop {
171        let (stream, peer) = listener.accept().await?;
172        info!(%peer, "ctp: accepted client connection");
173
174        // Cancel any existing connection before starting to serve the new one.
175        if let Some((task, token)) = connection_task.take() {
176            drop(token);
177            task.await;
178        }
179
180        let handler = handler_fn();
181        let version = version.clone();
182        let server_fqdn = server_fqdn.clone();
183        let metrics = metrics.clone();
184        let (cancel_tx, cancel_rx) = oneshot::channel();
185
186        let span = tracing::Span::current();
187        let handle = mz_ore::task::spawn(
188            || "ctp::connection",
189            async move {
190                let Err(error) = serve_connection(
191                    stream,
192                    handler,
193                    version,
194                    server_fqdn,
195                    idle_timeout,
196                    cancel_rx,
197                    metrics,
198                )
199                .await;
200                info!("ctp: connection failed: {error}");
201            }
202            .instrument(span),
203        );
204
205        connection_task = Some((handle, cancel_tx));
206    }
207}
208
209/// Serve a single CTP connection.
210async fn serve_connection<In, Out, H>(
211    stream: Stream,
212    mut handler: H,
213    version: Version,
214    server_fqdn: Option<String>,
215    timeout: Duration,
216    cancel_rx: oneshot::Receiver<()>,
217    metrics: impl Metrics<Out, In>,
218) -> anyhow::Result<Infallible>
219where
220    In: Message,
221    Out: Message,
222    H: GenericClient<In, Out>,
223{
224    let mut conn = Connection::start(stream, version, server_fqdn, timeout, metrics).await?;
225
226    let mut cancel_rx = cancel_rx;
227    loop {
228        tokio::select! {
229            // `Connection::recv` is documented to be cancel safe.
230            inbound = conn.recv() => {
231                let msg = inbound?;
232                handler.send(msg).await?;
233            },
234            // `GenericClient::recv` is documented to be cancel safe.
235            outbound = handler.recv() => match outbound? {
236                Some(msg) => conn.send(msg).await?,
237                None => bail!("client disconnected"),
238            },
239            _ = &mut cancel_rx => bail!("connection canceled"),
240        }
241    }
242}
243
244/// An active CTP connection.
245///
246/// This type encapsulates the core connection logic. It is used by both the client and the server
247/// implementation, with swapped `Out`/`In` types.
248///
249/// Each connection spawns two tasks:
250///
251///  * The send task is responsible for encoding and sending enqueued messages.
252///  * The recv task is responsible for receiving and decoding messages from the peer.
253///
254/// The separation into tasks provides some performance isolation between the sending and the
255/// receiving half of the connection.
256#[derive(Debug)]
257struct Connection<Out, In> {
258    /// Message sender connected to the send task.
259    msg_tx: mpsc::UnboundedSender<Out>,
260    /// Message receiver connected to the receive task.
261    msg_rx: mpsc::UnboundedReceiver<In>,
262    /// Receiver for errors encountered by connection tasks.
263    error_rx: ErrorRx,
264
265    /// Handles to connection tasks.
266    _tasks: [AbortOnDropHandle<()>; 2],
267}
268
269impl<Out: Message, In: Message> Connection<Out, In> {
270    /// The interval with which keepalives are emitted on idle connections.
271    const KEEPALIVE_INTERVAL: Duration = Duration::from_secs(1);
272    /// The minimum acceptable idle timeout.
273    ///
274    /// We want this to be significantly greater than `KEEPALIVE_INTERVAL`, to avoid connections
275    /// getting canceled unnecessarily.
276    const MIN_TIMEOUT: Duration = Duration::from_secs(2);
277
278    /// Start a new connection wrapping the given stream.
279    async fn start(
280        stream: Stream,
281        version: Version,
282        server_fqdn: Option<String>,
283        mut timeout: Duration,
284        metrics: impl Metrics<Out, In>,
285    ) -> anyhow::Result<Self> {
286        if timeout < Self::MIN_TIMEOUT {
287            warn!(
288                ?timeout,
289                "ctp: configured timeout is less than minimum timeout",
290            );
291            timeout = Self::MIN_TIMEOUT;
292        }
293
294        let (reader, writer) = stream.split();
295
296        // Apply the timeout to all connection reads and writes.
297        let reader = TimedReader::new(reader, timeout);
298        let writer = TimedWriter::new(writer, timeout);
299        // Track byte count metrics for all connection reads and writes.
300        let mut reader = metrics::Reader::new(reader, metrics.clone());
301        let mut writer = metrics::Writer::new(writer, metrics.clone());
302
303        handshake(&mut reader, &mut writer, version, server_fqdn).await?;
304
305        let (out_tx, out_rx) = mpsc::unbounded_channel();
306        let (in_tx, in_rx) = mpsc::unbounded_channel();
307        let (error_tx, error_rx) = error_channel();
308
309        let span = tracing::Span::current();
310        let send_task = mz_ore::task::spawn(
311            || "ctp::send",
312            Self::run_send_task(writer, out_rx, error_tx.clone(), metrics.clone())
313                .instrument(span.clone()),
314        );
315        let recv_task = mz_ore::task::spawn(
316            || "ctp::recv",
317            Self::run_recv_task(reader, in_tx, error_tx, metrics).instrument(span),
318        );
319
320        Ok(Self {
321            msg_tx: out_tx,
322            msg_rx: in_rx,
323            error_rx,
324            _tasks: [send_task.abort_on_drop(), recv_task.abort_on_drop()],
325        })
326    }
327
328    /// Enqueue a message for sending.
329    async fn send(&mut self, msg: Out) -> anyhow::Result<()> {
330        match self.msg_tx.send(msg) {
331            Ok(()) => Ok(()),
332            Err(_) => bail!(self.error_rx.collect().await),
333        }
334    }
335
336    /// Return a received message.
337    ///
338    /// # Cancel safety
339    ///
340    /// This method is cancel safe.
341    async fn recv(&mut self) -> anyhow::Result<In> {
342        // `mpcs::Receiver::recv` is documented to be cancel safe.
343        match self.msg_rx.recv().await {
344            Some(msg) => Ok(msg),
345            None => bail!(self.error_rx.collect().await),
346        }
347    }
348
349    /// Run a connection's send task.
350    async fn run_send_task<W: AsyncWrite + Unpin>(
351        mut writer: W,
352        mut msg_rx: mpsc::UnboundedReceiver<Out>,
353        error_tx: ErrorTx,
354        mut metrics: impl Metrics<Out, In>,
355    ) {
356        loop {
357            let msg = tokio::select! {
358                // `mpsc::UnboundedReceiver::recv` is cancel safe.
359                msg = msg_rx.recv() => match msg {
360                    Some(msg) => {
361                        trace!(?msg, "ctp: sending message");
362                        Some(msg)
363                    }
364                    None => break,
365                },
366                // `tokio::time::sleep` is cancel safe.
367                _ = tokio::time::sleep(Self::KEEPALIVE_INTERVAL) => {
368                    trace!("ctp: sending keepalive");
369                    None
370                },
371            };
372
373            if let Err(error) = write_message(&mut writer, msg.as_ref()).await {
374                warn!("ctp: send error: {error}");
375                error_tx.report(format!("send error: {error}"));
376                break;
377            };
378
379            if let Some(msg) = &msg {
380                metrics.message_sent(msg);
381            }
382        }
383    }
384
385    /// Run a connection's recv task.
386    async fn run_recv_task<R: AsyncRead + Unpin>(
387        mut reader: R,
388        msg_tx: mpsc::UnboundedSender<In>,
389        error_tx: ErrorTx,
390        mut metrics: impl Metrics<Out, In>,
391    ) {
392        loop {
393            match read_message(&mut reader).await {
394                Ok(msg) => {
395                    trace!(?msg, "ctp: received message");
396                    metrics.message_received(&msg);
397
398                    if msg_tx.send(msg).is_err() {
399                        break;
400                    }
401                }
402                Err(error) => {
403                    warn!("ctp: recv error: {error}");
404                    error_tx.report(format!("recv error: {error}"));
405                    break;
406                }
407            };
408        }
409    }
410}
411
412/// Error reported when a connection was closed without either task reporting an error.
413const CONNECTION_CLOSED: &str = "connection closed";
414
415/// Create a channel for reporting errors encountered by a connection's tasks.
416fn error_channel() -> (ErrorTx, ErrorRx) {
417    let (tx, rx) = watch::channel(None);
418    (ErrorTx(tx), ErrorRx(rx))
419}
420
421/// The sending half of a connection's error channel.
422#[derive(Clone, Debug)]
423struct ErrorTx(watch::Sender<Option<String>>);
424
425impl ErrorTx {
426    /// Report an error, unless an error was reported before.
427    fn report(&self, error: String) {
428        // The first error wins. A broken connection usually makes both tasks fail in sequence,
429        // with the first failure causing the second, so the first error is the one that explains
430        // what happened. For example, when the send task hits its idle deadline it drops its write
431        // half, whereupon the peer closes the connection and the recv task observes an EOF.
432        self.0.send_if_modified(|slot| match slot {
433            Some(_) => false,
434            None => {
435                *slot = Some(error);
436                true
437            }
438        });
439    }
440}
441
442/// The receiving half of a connection's error channel.
443#[derive(Debug)]
444struct ErrorRx(watch::Receiver<Option<String>>);
445
446impl ErrorRx {
447    /// Return the first error reported on this channel.
448    ///
449    /// If all [`ErrorTx`]s are dropped without an error being reported, returns
450    /// [`CONNECTION_CLOSED`]. Repeated calls return the same error.
451    async fn collect(&mut self) -> String {
452        // Wait for the first error to be reported, or for all connection tasks to shut down.
453        let _ = self.0.changed().await;
454        // Mark the current value as unseen, so the next `collect` call can return immediately.
455        self.0.mark_changed();
456
457        let error = self.0.borrow().clone();
458        error.unwrap_or_else(|| CONNECTION_CLOSED.into())
459    }
460}
461
462/// Perform the CTP handshake.
463///
464/// To perform the handshake, each endpoint sends the protocol magic number, followed by a
465/// `Hello` message. The `Hello` message contains information about the originating endpoint that
466/// is used by the receiver to validate compatibility with its peer. Only if both endpoints
467/// determine that they are compatible does the handshake succeed.
468async fn handshake<R, W>(
469    mut reader: R,
470    mut writer: W,
471    version: Version,
472    server_fqdn: Option<String>,
473) -> anyhow::Result<()>
474where
475    R: AsyncRead + Unpin,
476    W: AsyncWrite + Unpin,
477{
478    /// A randomly chosen magic number identifying CTP connections.
479    const MAGIC: u64 = 0x477574656e546167;
480
481    writer.write_u64(MAGIC).await?;
482
483    let hello = Hello {
484        version: version.clone(),
485        server_fqdn: server_fqdn.clone(),
486    };
487    write_message(&mut writer, Some(&hello)).await?;
488
489    let peer_magic = reader.read_u64().await?;
490    if peer_magic != MAGIC {
491        bail!("invalid protocol magic: {peer_magic:#x}");
492    }
493
494    let Hello {
495        version: peer_version,
496        server_fqdn: peer_server_fqdn,
497    } = read_message(&mut reader).await?;
498
499    if peer_version != version {
500        bail!("version mismatch: {peer_version} != {version}");
501    }
502    if let (Some(other), Some(mine)) = (&peer_server_fqdn, &server_fqdn) {
503        if other != mine {
504            bail!("server FQDN mismatch: {other} != {mine}");
505        }
506    }
507
508    Ok(())
509}
510
511/// A message for exchanging compatibility information during the CTP handshake.
512#[derive(Debug, Serialize, Deserialize)]
513struct Hello {
514    /// The version of the originating endpoint.
515    version: Version,
516    /// The FQDN of the server endpoint.
517    server_fqdn: Option<String>,
518}
519
520/// Write a message into the given writer.
521///
522/// The message can be `None`, in which case an empty message is written. This is used to implement
523/// keepalives. At the receiver, empty messages are ignored, but they do reset the read timeout.
524async fn write_message<W, M>(mut writer: W, msg: Option<&M>) -> anyhow::Result<()>
525where
526    W: AsyncWrite + Unpin,
527    M: Message,
528{
529    let bytes = match msg {
530        Some(msg) => &*wire_encode(msg)?,
531        None => &[],
532    };
533
534    let len = bytes.len().cast_into();
535    writer.write_u64(len).await?;
536    writer.write_all(bytes).await?;
537
538    Ok(())
539}
540
541/// Read a message from the given reader.
542async fn read_message<R, M>(mut reader: R) -> anyhow::Result<M>
543where
544    R: AsyncRead + Unpin,
545    M: Message,
546{
547    // Skip over any empty messages (i.e. keepalives).
548    let mut len = 0;
549    while len == 0 {
550        len = reader.read_u64().await?;
551    }
552
553    let mut bytes = vec![0; len.cast_into()];
554    reader.read_exact(&mut bytes).await?;
555
556    wire_decode(&bytes)
557}
558
559/// Encode a message for wire transport.
560fn wire_encode<M: Message>(msg: &M) -> anyhow::Result<Vec<u8>> {
561    let bytes = bincode::DefaultOptions::new().serialize(msg)?;
562    Ok(bytes)
563}
564
565/// Decode a wire frame back into a message.
566fn wire_decode<M: Message>(bytes: &[u8]) -> anyhow::Result<M> {
567    let msg = bincode::DefaultOptions::new().deserialize(bytes)?;
568    Ok(msg)
569}