1mod 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
49pub trait Message: Debug + Send + Sync + Serialize + DeserializeOwned + 'static {}
51impl<T: Debug + Send + Sync + Serialize + DeserializeOwned + 'static> Message for T {}
52
53#[derive(Debug)]
55pub struct Client<Out, In> {
56 conn: Connection<Out, In>,
57}
58
59impl<Out: Message, In: Message> Client<Out, In> {
60 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 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
86fn 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 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 async fn recv(&mut self) -> anyhow::Result<Option<In>> {
140 self.conn.recv().await.map(Some)
142 }
143}
144
145pub 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 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 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
209async 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 inbound = conn.recv() => {
231 let msg = inbound?;
232 handler.send(msg).await?;
233 },
234 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#[derive(Debug)]
257struct Connection<Out, In> {
258 msg_tx: mpsc::UnboundedSender<Out>,
260 msg_rx: mpsc::UnboundedReceiver<In>,
262 error_rx: ErrorRx,
264
265 _tasks: [AbortOnDropHandle<()>; 2],
267}
268
269impl<Out: Message, In: Message> Connection<Out, In> {
270 const KEEPALIVE_INTERVAL: Duration = Duration::from_secs(1);
272 const MIN_TIMEOUT: Duration = Duration::from_secs(2);
277
278 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 let reader = TimedReader::new(reader, timeout);
298 let writer = TimedWriter::new(writer, timeout);
299 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 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 async fn recv(&mut self) -> anyhow::Result<In> {
342 match self.msg_rx.recv().await {
344 Some(msg) => Ok(msg),
345 None => bail!(self.error_rx.collect().await),
346 }
347 }
348
349 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 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(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 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
412const CONNECTION_CLOSED: &str = "connection closed";
414
415fn error_channel() -> (ErrorTx, ErrorRx) {
417 let (tx, rx) = watch::channel(None);
418 (ErrorTx(tx), ErrorRx(rx))
419}
420
421#[derive(Clone, Debug)]
423struct ErrorTx(watch::Sender<Option<String>>);
424
425impl ErrorTx {
426 fn report(&self, error: String) {
428 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#[derive(Debug)]
444struct ErrorRx(watch::Receiver<Option<String>>);
445
446impl ErrorRx {
447 async fn collect(&mut self) -> String {
452 let _ = self.0.changed().await;
454 self.0.mark_changed();
456
457 let error = self.0.borrow().clone();
458 error.unwrap_or_else(|| CONNECTION_CLOSED.into())
459 }
460}
461
462async 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 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#[derive(Debug, Serialize, Deserialize)]
513struct Hello {
514 version: Version,
516 server_fqdn: Option<String>,
518}
519
520async 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
541async fn read_message<R, M>(mut reader: R) -> anyhow::Result<M>
543where
544 R: AsyncRead + Unpin,
545 M: Message,
546{
547 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
559fn wire_encode<M: Message>(msg: &M) -> anyhow::Result<Vec<u8>> {
561 let bytes = bincode::DefaultOptions::new().serialize(msg)?;
562 Ok(bytes)
563}
564
565fn wire_decode<M: Message>(bytes: &[u8]) -> anyhow::Result<M> {
567 let msg = bincode::DefaultOptions::new().deserialize(bytes)?;
568 Ok(msg)
569}