use crate::CanonicalMessage;
use anyhow::{anyhow, Result};
use bytes::{Bytes, BytesMut};
use futures::{FutureExt, SinkExt, StreamExt};
use tokio::io::{AsyncRead, AsyncWrite};
use tokio_util::codec::{Framed, LengthDelimitedCodec};
use tracing::warn;
const MAX_FRAME_BYTES: usize = 100 * 1024 * 1024;
const LENGTH_PREFIX_BYTES: usize = 4;
const STALLED_SEND_WARN_AFTER: std::time::Duration = std::time::Duration::from_secs(5);
pub(crate) type FramedIo<S> = Framed<S, LengthDelimitedCodec>;
pub(crate) fn wrap<S>(io: S) -> FramedIo<S>
where
S: AsyncRead + AsyncWrite,
{
LengthDelimitedCodec::builder()
.length_field_type::<u32>()
.max_frame_length(MAX_FRAME_BYTES)
.new_framed(io)
}
pub(crate) async fn send_batch<S>(
framed: &mut FramedIo<S>,
messages: &[CanonicalMessage],
peer: &str,
) -> Result<usize>
where
S: AsyncRead + AsyncWrite + Unpin + Send,
{
let data = rmp_serde::to_vec(messages)?;
if data.len() > MAX_FRAME_BYTES {
return Err(anyhow!(
"Encoded batch is {} bytes, exceeding the {} byte frame limit",
data.len(),
MAX_FRAME_BYTES
));
}
let len = data.len();
let mut send = std::pin::pin!(framed.send(Bytes::from(data)));
tokio::select! {
result = &mut send => result?,
_ = tokio::time::sleep(STALLED_SEND_WARN_AFTER) => {
warn!(
peer,
bytes = len,
count = messages.len(),
stalled_for_secs = STALLED_SEND_WARN_AFTER.as_secs(),
"IPC send is blocked: the batch is larger than the socket buffer and the \
consumer is not draining it. Check that the consumer side is running and \
reading. Still waiting."
);
send.await?
}
}
Ok(len)
}
pub(crate) async fn recv_batch<S>(framed: &mut FramedIo<S>) -> Result<Vec<CanonicalMessage>>
where
S: AsyncRead + AsyncWrite + Unpin + Send,
{
match framed.next().await {
Some(frame) => decode(&frame?),
None => Err(std::io::Error::from(std::io::ErrorKind::UnexpectedEof).into()),
}
}
pub(crate) fn try_recv_batch<S>(framed: &mut FramedIo<S>) -> Result<Option<Vec<CanonicalMessage>>>
where
S: AsyncRead + AsyncWrite + Unpin + Send,
{
match framed.next().now_or_never() {
Some(Some(frame)) => decode(&frame?).map(Some),
Some(None) => Err(std::io::Error::from(std::io::ErrorKind::UnexpectedEof).into()),
None => Ok(None),
}
}
pub(crate) fn buffered_frames<S>(framed: &FramedIo<S>) -> usize {
let buf: &BytesMut = framed.read_buffer();
let mut offset = 0usize;
let mut frames = 0usize;
while buf.len() >= offset + LENGTH_PREFIX_BYTES {
let mut prefix = [0u8; LENGTH_PREFIX_BYTES];
prefix.copy_from_slice(&buf[offset..offset + LENGTH_PREFIX_BYTES]);
let payload = u32::from_be_bytes(prefix) as usize;
let end = match offset
.checked_add(LENGTH_PREFIX_BYTES)
.and_then(|o| o.checked_add(payload))
{
Some(end) => end,
None => break,
};
if buf.len() < end {
break;
}
offset = end;
frames += 1;
}
frames
}
fn decode(frame: &[u8]) -> Result<Vec<CanonicalMessage>> {
rmp_serde::from_slice(frame)
.map_err(|e| anyhow!("Failed to decode IPC batch ({} bytes): {}", frame.len(), e))
}
#[cfg(test)]
mod tests {
use super::*;
use tokio::io::duplex;
#[tokio::test]
async fn roundtrip_preserves_batch() {
let (a, b) = duplex(64 * 1024);
let mut writer = wrap(a);
let mut reader = wrap(b);
let sent = vec![
CanonicalMessage::from_vec(b"one"),
CanonicalMessage::from_vec(b"two"),
];
send_batch(&mut writer, &sent, "test").await.unwrap();
let received = recv_batch(&mut reader).await.unwrap();
assert_eq!(received.len(), 2);
assert_eq!(received[0].payload.as_ref(), b"one");
assert_eq!(received[1].payload.as_ref(), b"two");
}
#[tokio::test]
async fn try_recv_returns_none_when_idle() {
let (a, b) = duplex(64 * 1024);
let _writer = wrap(a);
let mut reader = wrap(b);
assert!(try_recv_batch(&mut reader).unwrap().is_none());
}
#[tokio::test]
async fn try_recv_yields_a_buffered_frame() {
let (a, b) = duplex(64 * 1024);
let mut writer = wrap(a);
let mut reader = wrap(b);
send_batch(&mut writer, &[CanonicalMessage::from_vec(b"ready")], "test")
.await
.unwrap();
let received = try_recv_batch(&mut reader)
.unwrap()
.expect("a frame is already on the wire");
assert_eq!(received[0].payload.as_ref(), b"ready");
}
#[tokio::test]
async fn clean_eof_surfaces_as_unexpected_eof() {
let (a, b) = duplex(64 * 1024);
drop(a);
let mut reader = wrap(b);
let err = recv_batch(&mut reader).await.unwrap_err();
let io_err = err
.downcast_ref::<std::io::Error>()
.expect("EOF should surface as an io error");
assert_eq!(io_err.kind(), std::io::ErrorKind::UnexpectedEof);
}
#[tokio::test]
async fn cancelled_read_resumes_without_desync() {
let (a, b) = duplex(64 * 1024);
let mut writer = wrap(a);
let mut reader = wrap(b);
let messages = vec![CanonicalMessage::from_vec(b"survives-cancellation")];
let encoded = rmp_serde::to_vec(&messages).unwrap();
{
use tokio::io::AsyncWriteExt;
let io = writer.get_mut();
io.write_all(&(encoded.len() as u32).to_be_bytes())
.await
.unwrap();
io.write_all(&encoded[..4]).await.unwrap();
io.flush().await.unwrap();
}
tokio::select! {
_ = recv_batch(&mut reader) => panic!("frame is incomplete, read must not finish"),
_ = tokio::time::sleep(std::time::Duration::from_millis(50)) => {}
}
{
use tokio::io::AsyncWriteExt;
let io = writer.get_mut();
io.write_all(&encoded[4..]).await.unwrap();
io.flush().await.unwrap();
}
let received = recv_batch(&mut reader).await.unwrap();
assert_eq!(received[0].payload.as_ref(), b"survives-cancellation");
}
#[tokio::test]
async fn buffered_frames_counts_whole_frames_only() {
let (a, b) = duplex(64 * 1024);
let mut writer = wrap(a);
let mut reader = wrap(b);
send_batch(&mut writer, &[CanonicalMessage::from_vec(b"a")], "test")
.await
.unwrap();
send_batch(&mut writer, &[CanonicalMessage::from_vec(b"b")], "test")
.await
.unwrap();
let _ = recv_batch(&mut reader).await.unwrap();
assert_eq!(buffered_frames(&reader), 1);
let _ = recv_batch(&mut reader).await.unwrap();
assert_eq!(buffered_frames(&reader), 0);
}
}