use bytes::BytesMut;
use crate::error::ZmqError;
use crate::protocol::zmtp::actions::EngineOutput;
use crate::protocol::zmtp::engine::ZmtpEngine;
use crate::transport::ZmtpReadHalf;
use super::INGRESS_GREEDY_CHUNK;
#[cfg(not(target_os = "macos"))]
const READ_CHUNK: usize = 65536 * 2;
#[cfg(target_os = "macos")]
const READ_CHUNK: usize = 65536;
const FRAME_READ_CHUNK: usize = 16 * 1024;
pub(crate) struct ZmqMessageProcessor;
impl ZmqMessageProcessor {
pub(crate) fn new() -> Self {
Self
}
pub(crate) async fn read_and_process<RH: ZmtpReadHalf>(
&mut self,
reader: &mut RH,
engine: &mut ZmtpEngine,
frame_in_place: bool,
) -> Result<EngineOutput, ZmqError> {
if engine.buffer_len() > 16 * 1024 * 1024 {
return Err(ZmqError::ResourceLimitReached);
}
if frame_in_place {
self.read_frame_in_place(reader, engine).await
} else {
self.read_accumulate(reader, engine).await
}
}
async fn read_accumulate<RH: ZmtpReadHalf>(
&mut self,
reader: &mut RH,
engine: &mut ZmtpEngine,
) -> Result<EngineOutput, ZmqError> {
use tokio::io::AsyncReadExt;
let mut buf = BytesMut::with_capacity(INGRESS_GREEDY_CHUNK);
let n = reader
.read_buf(&mut buf)
.await
.map_err(|e| ZmqError::from_io_endpoint(e, "ingress read"))?;
if n == 0 {
return Err(ZmqError::ConnectionClosed);
}
let mut total_read = n;
let max_greedy_read = engine.config().rcvbatch_bytes.max(INGRESS_GREEDY_CHUNK);
let mut greedy_buf = [0u8; INGRESS_GREEDY_CHUNK];
while total_read < max_greedy_read {
match reader.try_read_chunk(&mut greedy_buf) {
Ok(0) => break,
Ok(k) => {
buf.extend_from_slice(&greedy_buf[..k]);
total_read += k;
}
Err(e) if e.kind() == std::io::ErrorKind::WouldBlock => break,
Err(e) => return Err(ZmqError::from_io_endpoint(e, "ingress greedy read")),
}
}
Ok(engine.on_network_bytes(buf.freeze()))
}
async fn read_frame_in_place<RH: ZmtpReadHalf>(
&mut self,
reader: &mut RH,
engine: &mut ZmtpEngine,
) -> Result<EngineOutput, ZmqError> {
use tokio::io::AsyncReadExt;
let max_greedy_read = engine.config().rcvbatch_bytes.max(INGRESS_GREEDY_CHUNK);
let carry = engine.carry();
let mut buf = BytesMut::with_capacity(carry.len() + FRAME_READ_CHUNK);
buf.extend_from_slice(carry);
let n = reader
.read_buf(&mut buf)
.await
.map_err(|e| ZmqError::from_io_endpoint(e, "ingress read"))?;
if n == 0 {
return Err(ZmqError::ConnectionClosed);
}
let mut total_read = n;
while total_read < max_greedy_read {
if buf.capacity() - buf.len() < 4096 {
buf.reserve(FRAME_READ_CHUNK);
}
match reader.try_read_buf(&mut buf) {
Ok(0) => break,
Ok(k) => {
total_read += k;
}
Err(e) if e.kind() == std::io::ErrorKind::WouldBlock => break,
Err(e) => return Err(ZmqError::from_io_endpoint(e, "ingress greedy read")),
}
}
Ok(engine.drive_buf(buf))
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::protocol::zmtp::engine::ZmtpEngine;
use crate::socket::options::ZmtpEngineConfig;
use std::io;
use std::pin::Pin;
use std::sync::Arc;
use std::task::{Context, Poll};
use tokio::io::{AsyncRead, ReadBuf};
#[derive(Debug)]
struct InfiniteMockReader;
impl AsyncRead for InfiniteMockReader {
fn poll_read(
self: Pin<&mut Self>,
_cx: &mut Context<'_>,
buf: &mut ReadBuf<'_>,
) -> Poll<io::Result<()>> {
let dummy = vec![0u8; buf.remaining()];
buf.put_slice(&dummy);
Poll::Ready(Ok(()))
}
}
impl crate::transport::ZmtpReadHalf for InfiniteMockReader {
fn try_read_chunk(&mut self, buf: &mut [u8]) -> io::Result<usize> {
for b in buf.iter_mut() {
*b = 0;
}
Ok(buf.len())
}
fn try_read_buf(&mut self, buf: &mut BytesMut) -> io::Result<usize> {
let spare = buf.capacity() - buf.len();
if spare == 0 {
return Ok(0);
}
buf.extend_from_slice(&vec![0u8; spare]);
Ok(spare)
}
}
#[tokio::test]
async fn test_mre_ingress_starvation_deadlock() {
let mut processor = ZmqMessageProcessor::new();
let config = Arc::new(ZmtpEngineConfig::default());
let mut engine = ZmtpEngine::new(true, config);
let mut mock_reader = InfiniteMockReader;
let result = tokio::time::timeout(
std::time::Duration::from_millis(50),
processor.read_and_process(&mut mock_reader, &mut engine, false),
)
.await;
assert!(
result.is_ok(),
"REGRESSION: ZmaqMessageProcessor is trapped in an infinite greedy-read loop!"
);
}
}