use super::{Candidate, RequestRegistration, ResponseIdentity, Transport, extract_request_id};
use crate::error::{ConstructionStage, Error, Result};
use crate::message_size::ReceiveLimits;
use bytes::{Bytes, BytesMut};
use std::net::SocketAddr;
use std::sync::Arc;
use std::sync::atomic::{AtomicBool, Ordering};
use std::time::Duration;
use tokio::io::{AsyncReadExt, AsyncWriteExt};
use tokio::net::TcpStream;
use tokio::sync::{Mutex, OwnedMutexGuard};
#[cfg(test)]
use tokio::time::timeout;
const DEFAULT_MAX_MESSAGE_SIZE: usize = 10 * 1024 * 1024;
#[derive(Debug, Clone)]
#[non_exhaustive]
pub struct TcpOptions {
pub max_message_size: usize,
pub send_capacity: usize,
}
impl Default for TcpOptions {
fn default() -> Self {
Self {
max_message_size: DEFAULT_MAX_MESSAGE_SIZE,
send_capacity: DEFAULT_MAX_MESSAGE_SIZE,
}
}
}
#[derive(Debug)]
pub struct TcpTransportBuilder {
connect_timeout: Option<Duration>,
options: TcpOptions,
}
impl TcpTransportBuilder {
#[must_use]
pub fn new() -> Self {
Self {
connect_timeout: None,
options: TcpOptions::default(),
}
}
#[must_use]
pub fn connect_timeout(mut self, timeout: Duration) -> Self {
self.connect_timeout = Some(timeout);
self
}
#[must_use]
pub fn max_message_size(mut self, size: usize) -> Self {
self.options.max_message_size = size;
self
}
#[must_use]
pub fn send_capacity(mut self, size: usize) -> Self {
self.options.send_capacity = size;
self
}
pub async fn connect(self, target: SocketAddr) -> Result<TcpTransport> {
let started = tokio::time::Instant::now();
let connect_deadline = self
.connect_timeout
.map(|timeout| tcp_deadline(timeout, "TCP connect timeout"))
.transpose()?;
let receive_limits = ReceiveLimits::tcp(self.options.max_message_size)
.map_err(|error| Error::Config(error.to_string().into()).boxed())?;
let stream = match (self.connect_timeout, connect_deadline) {
(Some(_), Some(deadline)) if tokio::time::Instant::now() >= deadline => {
return Err(Error::ConstructionTimeout {
target: target.into(),
stage: ConstructionStage::Connect,
elapsed: started.elapsed(),
}
.boxed());
}
(Some(_), Some(deadline)) => {
tokio::time::timeout_at(deadline, TcpStream::connect(target))
.await
.map_err(|_| {
Error::ConstructionTimeout {
target: target.into(),
stage: ConstructionStage::Connect,
elapsed: started.elapsed(),
}
.boxed()
})?
.map_err(|e| Error::Network { target, source: e }.boxed())?
}
(None, None) => TcpStream::connect(target)
.await
.map_err(|e| Error::Network { target, source: e }.boxed())?,
_ => unreachable!("deadline presence follows timeout presence"),
};
let local_addr = stream
.local_addr()
.map_err(|e| Error::Network { target, source: e }.boxed())?;
Ok(TcpTransport {
inner: Arc::new(TcpTransportInner {
stream: Arc::new(Mutex::new(stream)),
target,
local_addr,
receive_limits,
send_capacity: self.options.send_capacity,
poisoned: AtomicBool::new(false),
}),
})
}
}
impl Default for TcpTransportBuilder {
fn default() -> Self {
Self::new()
}
}
#[derive(Clone)]
pub struct TcpTransport {
inner: Arc<TcpTransportInner>,
}
struct TcpTransportInner {
stream: Arc<Mutex<TcpStream>>,
target: SocketAddr,
local_addr: SocketAddr,
receive_limits: ReceiveLimits,
send_capacity: usize,
poisoned: AtomicBool,
}
impl TcpTransportInner {
fn is_poisoned(&self) -> bool {
self.poisoned.load(Ordering::Acquire)
}
}
struct TcpTransactionGuard<'a> {
poisoned: &'a AtomicBool,
armed: bool,
}
impl<'a> TcpTransactionGuard<'a> {
fn new(inner: &'a TcpTransportInner) -> Self {
Self {
poisoned: &inner.poisoned,
armed: true,
}
}
fn unarmed(inner: &'a TcpTransportInner) -> Self {
Self {
poisoned: &inner.poisoned,
armed: false,
}
}
#[cfg(test)]
fn for_test(poisoned: &'a AtomicBool) -> Self {
Self {
poisoned,
armed: true,
}
}
fn disarm(&mut self) {
self.armed = false;
}
fn arm(&mut self) {
self.armed = true;
}
}
impl Drop for TcpTransactionGuard<'_> {
fn drop(&mut self) {
if self.armed {
self.poisoned.store(true, Ordering::Release);
}
}
}
impl TcpTransport {
pub async fn connect(target: SocketAddr) -> Result<Self> {
Self::builder().connect(target).await
}
pub async fn connect_timeout(target: SocketAddr, connect_timeout: Duration) -> Result<Self> {
Self::builder()
.connect_timeout(connect_timeout)
.connect(target)
.await
}
#[must_use]
pub fn builder() -> TcpTransportBuilder {
TcpTransportBuilder::new()
}
pub async fn from_socket(
socket: tokio::net::TcpSocket,
target: SocketAddr,
options: TcpOptions,
) -> Result<Self> {
let receive_limits = ReceiveLimits::tcp(options.max_message_size)
.map_err(|error| Error::Config(error.to_string().into()).boxed())?;
let stream = socket
.connect(target)
.await
.map_err(|e| Error::Network { target, source: e }.boxed())?;
let local_addr = stream
.local_addr()
.map_err(|e| Error::Network { target, source: e }.boxed())?;
Ok(Self {
inner: Arc::new(TcpTransportInner {
stream: Arc::new(Mutex::new(stream)),
target,
local_addr,
receive_limits,
send_capacity: options.send_capacity,
poisoned: AtomicBool::new(false),
}),
})
}
}
impl Transport for TcpTransport {
async fn send(&self, data: &[u8]) -> Result<()> {
crate::message_size::enforce_outbound_size(data.len(), self.send_capacity())?;
let mut stream = self.inner.stream.clone().lock_owned().await;
let target = self.inner.target;
if self.inner.is_poisoned() {
return Err(Error::Closed { target }.boxed());
}
let mut transaction = TcpTransactionGuard::new(&self.inner);
write_message(&mut *stream, target, data).await?;
transaction.disarm();
Ok(())
}
async fn send_with_timeout(&self, data: &[u8], timeout: Duration) -> Result<()> {
crate::message_size::enforce_outbound_size(data.len(), self.send_capacity())?;
let target = self.inner.target;
let deadline = transaction_deadline(timeout)?;
let mut stream = lock_stream_before(&self.inner, deadline, timeout).await?;
if self.inner.is_poisoned() {
return Err(Error::Closed { target }.boxed());
}
let mut transaction = TcpTransactionGuard::unarmed(&self.inner);
tokio::time::timeout_at(
deadline,
arm_when_polled(&mut transaction, write_message(&mut *stream, target, data)),
)
.await
.map_err(|_| timeout_error(target, timeout))??;
transaction.disarm();
Ok(())
}
async fn request_with<T, F>(
&self,
data: &[u8],
registration: RequestRegistration,
validate: F,
) -> Result<T>
where
T: Send,
F: FnMut(Bytes, SocketAddr) -> Result<Candidate<T>> + Send,
{
crate::message_size::enforce_outbound_size(data.len(), self.send_capacity())?;
let request_id = registration.request_id();
let started = tokio::time::Instant::now();
let deadline = registration.deadline();
let recv_timeout = deadline.saturating_duration_since(started);
let target = self.inner.target;
let mut stream = lock_stream_before(&self.inner, deadline, recv_timeout).await?;
if self.inner.is_poisoned() {
return Err(Error::Closed { target }.boxed());
}
let mut transaction = TcpTransactionGuard::unarmed(&self.inner);
tokio::time::timeout_at(
deadline,
arm_when_polled(&mut transaction, write_message(&mut *stream, target, data)),
)
.await
.map_err(|_| timeout_error(target, recv_timeout))??;
let result = tokio::time::timeout_at(
deadline,
arm_when_polled(
&mut transaction,
read_validated_message(
&mut stream,
target,
self.inner.receive_limits.accepted(),
®istration,
validate,
),
),
)
.await
.map_err(|_| CorrelatedReadError::Timeout)
.and_then(std::convert::identity);
let value = finish_correlated_read(result, target, request_id, recv_timeout)?;
transaction.disarm();
Ok(value)
}
fn peer_addr(&self) -> SocketAddr {
self.inner.target
}
fn local_addr(&self) -> SocketAddr {
self.inner.local_addr
}
fn is_reliable(&self) -> bool {
true
}
fn receive_limits(&self) -> ReceiveLimits {
self.inner.receive_limits
}
fn send_capacity(&self) -> usize {
self.inner.send_capacity
}
}
#[cfg(test)]
impl TcpTransport {
async fn recv(&self, registration: RequestRegistration) -> Result<(Bytes, SocketAddr)> {
let request_id = registration.request_id();
let started = tokio::time::Instant::now();
let deadline = registration.deadline();
let elapsed = deadline.saturating_duration_since(started);
let target = self.inner.target;
let mut stream = lock_stream_before(&self.inner, deadline, elapsed).await?;
if self.inner.is_poisoned() {
return Err(Error::Closed { target }.boxed());
}
let mut transaction = TcpTransactionGuard::unarmed(&self.inner);
let result = tokio::time::timeout_at(
deadline,
arm_when_polled(
&mut transaction,
read_validated_message(
&mut stream,
target,
self.inner.receive_limits.accepted(),
®istration,
|data, source| Ok(Candidate::Accept((data, source))),
),
),
)
.await
.map_err(|_| CorrelatedReadError::Timeout)
.and_then(std::convert::identity);
let value = finish_correlated_read(result, target, request_id, elapsed)?;
transaction.disarm();
Ok(value)
}
async fn request(
&self,
data: &[u8],
registration: RequestRegistration,
) -> Result<(Bytes, SocketAddr)> {
self.request_with(data, registration, |response, source| {
Ok(Candidate::Accept((response, source)))
})
.await
}
}
enum CorrelatedReadError {
Framing(Box<Error>),
Validation(Box<Error>),
Timeout,
}
fn tcp_deadline(timeout: Duration, description: &str) -> Result<tokio::time::Instant> {
tokio::time::Instant::now()
.checked_add(timeout)
.ok_or_else(|| {
Error::Config(format!("{description} exceeds the representable deadline").into())
.boxed()
})
}
fn transaction_deadline(timeout: Duration) -> Result<tokio::time::Instant> {
tcp_deadline(timeout, "TCP timeout")
}
async fn lock_stream_before(
inner: &TcpTransportInner,
deadline: tokio::time::Instant,
timeout: Duration,
) -> Result<OwnedMutexGuard<TcpStream>> {
let target = inner.target;
if tokio::time::Instant::now() >= deadline {
return Err(timeout_error(target, timeout));
}
let stream = tokio::time::timeout_at(deadline, inner.stream.clone().lock_owned())
.await
.map_err(|_| timeout_error(target, timeout))?;
if tokio::time::Instant::now() >= deadline {
return Err(timeout_error(target, timeout));
}
Ok(stream)
}
fn timeout_error(target: SocketAddr, elapsed: Duration) -> Box<Error> {
Error::Timeout {
target,
elapsed,
retries: 0,
}
.boxed()
}
async fn arm_when_polled<F>(transaction: &mut TcpTransactionGuard<'_>, future: F) -> F::Output
where
F: std::future::Future,
{
tokio::pin!(future);
std::future::poll_fn(|context| {
transaction.arm();
future.as_mut().poll(context)
})
.await
}
async fn write_message<W>(stream: &mut W, target: SocketAddr, data: &[u8]) -> Result<()>
where
W: tokio::io::AsyncWrite + Unpin,
{
stream
.write_all(data)
.await
.map_err(|source| Error::Network { target, source }.boxed())?;
stream
.flush()
.await
.map_err(|source| Error::Network { target, source }.boxed())
}
async fn read_validated_message<T, F>(
stream: &mut TcpStream,
target: SocketAddr,
max_message_size: usize,
registration: &RequestRegistration,
mut validate: F,
) -> std::result::Result<T, CorrelatedReadError>
where
F: FnMut(Bytes, SocketAddr) -> Result<Candidate<T>>,
{
let request_id = registration.request_id();
loop {
let frame = read_ber_message(stream, target, max_message_size)
.await
.map_err(CorrelatedReadError::Framing)?;
let Some(frame_id) = extract_request_id(&frame) else {
tracing::debug!(target: "async_snmp::transport::tcp", { request_id, %target }, "complete response frame has no extractable correlation ID");
continue;
};
if frame_id != request_id && !registration.aliases().contains(&frame_id) {
tracing::debug!(target: "async_snmp::transport::tcp", { request_id, frame_id, %target }, "stale response frame skipped");
continue;
}
match registration.evaluate_response_identity(&frame, true) {
ResponseIdentity::Match => {}
ResponseIdentity::AcceptedCommunityMismatch => {
tracing::warn!(target: "async_snmp::transport::tcp", { request_id, %target }, "accepted rewritten response community");
}
ResponseIdentity::Reject => {
tracing::debug!(target: "async_snmp::transport::tcp", { request_id, %target }, "response rejected by registered identity correlation");
continue;
}
}
match validate(frame, target).map_err(CorrelatedReadError::Validation)? {
Candidate::Accept(value) => return Ok(value),
Candidate::Reject => {
tracing::debug!(target: "async_snmp::transport::tcp", { request_id, %target }, "response rejected by client validation");
}
}
}
}
fn finish_correlated_read<T>(
result: std::result::Result<T, CorrelatedReadError>,
target: SocketAddr,
request_id: i32,
recv_timeout: Duration,
) -> Result<T> {
match result {
Ok(value) => Ok(value),
Err(CorrelatedReadError::Framing(error)) | Err(CorrelatedReadError::Validation(error)) => {
Err(error)
}
Err(CorrelatedReadError::Timeout) => {
tracing::debug!(target: "async_snmp::transport::tcp", { request_id, %target, elapsed = ?recv_timeout }, "transport timeout");
Err(timeout_error(target, recv_timeout))
}
}
}
async fn read_ber_message(
stream: &mut TcpStream,
target: SocketAddr,
max_message_size: usize,
) -> Result<Bytes> {
let mut tag_buf = [0u8; 1];
stream
.read_exact(&mut tag_buf)
.await
.map_err(|e| Error::Network { target, source: e }.boxed())?;
let tag = tag_buf[0];
if tag & 0x1f == 0x1f {
tracing::debug!(target: "async_snmp::transport::tcp", { actual_tag = tag, %target }, "multi-byte tag not supported");
return Err(Error::Decode(
crate::DecodeError::new(
0,
crate::DecodeErrorKind::UnsupportedMultiOctetTag { first_octet: tag },
)
.with_peer(target),
)
.boxed());
}
if tag != 0x30 {
tracing::debug!(target: "async_snmp::transport::tcp", { expected_tag = 0x30, actual_tag = tag, %target }, "invalid SNMP message tag");
return Err(Error::Decode(
crate::DecodeError::new(
0,
crate::DecodeErrorKind::UnexpectedTag {
expected: 0x30,
actual: tag,
},
)
.with_peer(target),
)
.boxed());
}
let mut first_len_byte = [0u8; 1];
stream
.read_exact(&mut first_len_byte)
.await
.map_err(|e| Error::Network { target, source: e }.boxed())?;
let (content_len, len_bytes) = match first_len_byte[0].cmp(&0x80) {
std::cmp::Ordering::Less => {
(first_len_byte[0] as usize, vec![first_len_byte[0]])
}
std::cmp::Ordering::Equal => {
tracing::debug!(target: "async_snmp::transport::tcp", { %target }, "indefinite length encoding not supported");
return Err(Error::Decode(
crate::DecodeError::new(1, crate::DecodeErrorKind::IndefiniteLength)
.with_peer(target),
)
.boxed());
}
std::cmp::Ordering::Greater => {
let num_len_bytes = (first_len_byte[0] & 0x7F) as usize;
if num_len_bytes > 4 {
tracing::debug!(target: "async_snmp::transport::tcp", { octets = num_len_bytes, %target }, "length encoding too long");
return Err(Error::Decode(
crate::DecodeError::new(
1,
crate::DecodeErrorKind::LengthTooLong {
octets: num_len_bytes,
},
)
.with_peer(target),
)
.boxed());
}
let mut len_bytes_buf = vec![0u8; num_len_bytes];
stream
.read_exact(&mut len_bytes_buf)
.await
.map_err(|e| Error::Network { target, source: e }.boxed())?;
let mut length: usize = 0;
for &b in &len_bytes_buf {
length = (length << 8) | (b as usize);
}
let mut all_len_bytes = vec![first_len_byte[0]];
all_len_bytes.extend_from_slice(&len_bytes_buf);
(length, all_len_bytes)
}
};
let total_len = 1usize
.checked_add(len_bytes.len())
.and_then(|header_len| header_len.checked_add(content_len))
.ok_or_else(|| {
Error::Decode(
crate::DecodeError::new(1, crate::DecodeErrorKind::IntegerOverflow)
.with_peer(target),
)
.boxed()
})?;
if total_len > max_message_size {
tracing::warn!(target: "async_snmp::transport::tcp", { size = total_len, max = max_message_size, %target }, "message size exceeds limit");
return Err(Error::Decode(
crate::DecodeError::new(
0,
crate::DecodeErrorKind::MessageTooLarge {
size: total_len,
maximum: max_message_size,
},
)
.with_peer(target),
)
.boxed());
}
let mut content = vec![0u8; content_len];
stream
.read_exact(&mut content)
.await
.map_err(|e| Error::Network { target, source: e }.boxed())?;
let mut message = BytesMut::with_capacity(total_len);
message.extend_from_slice(&[tag]);
message.extend_from_slice(&len_bytes);
message.extend_from_slice(&content);
Ok(message.freeze())
}
#[cfg(test)]
mod tests {
use super::*;
use std::io;
use std::pin::Pin;
use std::task::{Context, Poll};
use tokio::io::{AsyncWrite, AsyncWriteExt};
use tokio::net::TcpListener;
fn deadline_after(timeout: Duration) -> tokio::time::Instant {
tokio::time::Instant::now() + timeout
}
enum WriteFailure {
Write,
Flush,
}
struct FailingWriter(WriteFailure);
impl AsyncWrite for FailingWriter {
fn poll_write(
self: Pin<&mut Self>,
_cx: &mut Context<'_>,
buf: &[u8],
) -> Poll<io::Result<usize>> {
match self.0 {
WriteFailure::Write => Poll::Ready(Err(io::Error::other("write failed"))),
WriteFailure::Flush => Poll::Ready(Ok(buf.len())),
}
}
fn poll_flush(self: Pin<&mut Self>, _cx: &mut Context<'_>) -> Poll<io::Result<()>> {
match self.0 {
WriteFailure::Write => Poll::Ready(Ok(())),
WriteFailure::Flush => Poll::Ready(Err(io::Error::other("flush failed"))),
}
}
fn poll_shutdown(self: Pin<&mut Self>, _cx: &mut Context<'_>) -> Poll<io::Result<()>> {
Poll::Ready(Ok(()))
}
}
#[tokio::test]
async fn guarded_write_and_flush_failures_poison() {
let target = "127.0.0.1:161".parse().unwrap();
for failure in [WriteFailure::Write, WriteFailure::Flush] {
let poisoned = AtomicBool::new(false);
let mut writer = FailingWriter(failure);
{
let _transaction = TcpTransactionGuard::for_test(&poisoned);
let error = write_message(&mut writer, target, b"message")
.await
.unwrap_err();
assert!(matches!(*error, Error::Network { .. }));
}
assert!(poisoned.load(Ordering::Acquire));
}
}
#[tokio::test]
async fn unrepresentable_connect_deadline_starts_no_connection() {
let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
let server_addr = listener.local_addr().unwrap();
let error = TcpTransport::connect_timeout(server_addr, Duration::MAX)
.await
.err()
.expect("unrepresentable connect deadline must fail");
assert!(matches!(*error, Error::Config(_)));
assert!(
tokio::time::timeout(Duration::from_millis(50), listener.accept())
.await
.is_err(),
"invalid connect deadline started stream I/O"
);
}
#[tokio::test]
async fn test_tcp_send_recv() {
let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
let server_addr = listener.local_addr().unwrap();
let server = tokio::spawn(async move {
let (mut socket, _) = listener.accept().await.unwrap();
let mut buf = vec![0u8; 1024];
let n = socket.read(&mut buf).await.unwrap();
let response = [
0x30, 0x1c, 0x02, 0x01, 0x01, 0x04, 0x06, 0x70, 0x75, 0x62, 0x6c, 0x69, 0x63, 0xa2, 0x0f, 0x02, 0x01, 0x01, 0x02, 0x01, 0x00, 0x02, 0x01, 0x00, 0x30, 0x04, 0x30, 0x02, 0x05, 0x00, ];
socket.write_all(&response).await.unwrap();
n
});
let transport = TcpTransport::connect(server_addr).await.unwrap();
let request = [
0x30, 0x1a, 0x02, 0x01, 0x01, 0x04, 0x06, 0x70, 0x75, 0x62, 0x6c, 0x69, 0x63, 0xa0, 0x0d, 0x02, 0x01, 0x01, 0x02, 0x01, 0x00, 0x02, 0x01, 0x00, 0x30, 0x02, 0x30, 0x00,
];
transport.send(&request).await.unwrap();
let registration = RequestRegistration::test_unchecked(1, Duration::from_secs(5));
let (response, source) = transport.recv(registration).await.unwrap();
assert_eq!(source, server_addr);
assert_eq!(response[0], 0x30); assert!(response.len() > 10);
server.await.unwrap();
}
#[tokio::test]
async fn test_tcp_long_length_form() {
let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
let server_addr = listener.local_addr().unwrap();
let server = tokio::spawn(async move {
let (mut socket, _) = listener.accept().await.unwrap();
let mut request = [0u8; 31];
socket.read_exact(&mut request).await.unwrap();
let response = build_response_with_id(1);
let mut long_form = vec![0x30, 0x81, response[1]];
long_form.extend_from_slice(&response[2..]);
socket.write_all(&long_form).await.unwrap();
});
let transport = TcpTransport::connect(server_addr).await.unwrap();
let registration = RequestRegistration::test_unchecked(1, Duration::from_secs(5));
let (response, _) = transport
.request(&build_request_with_id(1), registration)
.await
.unwrap();
assert_eq!(response.len(), 32);
assert_eq!(&response[..3], &[0x30, 0x81, 0x1d]);
server.await.unwrap();
}
#[tokio::test]
async fn test_tcp_advertised_max_matches_accepted_limit() {
let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
let server_addr = listener.local_addr().unwrap();
tokio::spawn(async move {
let mut conns = Vec::new();
while let Ok((socket, _)) = listener.accept().await {
conns.push(socket);
}
});
let transport = TcpTransport::connect(server_addr).await.unwrap();
assert_eq!(
transport.receive_limits().advertised().as_usize(),
transport.inner.receive_limits.accepted(),
"advertised msgMaxSize must equal the accepted total-message limit"
);
assert_eq!(
transport.receive_limits().advertised().as_usize(),
DEFAULT_MAX_MESSAGE_SIZE
);
let custom = 512 * 1024;
let transport = TcpTransport::builder()
.max_message_size(custom)
.connect(server_addr)
.await
.unwrap();
assert_eq!(transport.receive_limits().advertised().as_usize(), custom);
assert_eq!(
transport.receive_limits().advertised().as_usize(),
transport.inner.receive_limits.accepted()
);
}
#[tokio::test]
async fn receive_advertisement_and_send_capacity_are_independent() {
let request = build_request_with_id(1);
let exact_limit = request.len();
let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
let server_addr = listener.local_addr().unwrap();
let (inspect_tx, inspect_rx) = tokio::sync::oneshot::channel();
let (inspected_tx, inspected_rx) = tokio::sync::oneshot::channel();
let server = tokio::spawn(async move {
let (mut socket, _) = listener.accept().await.unwrap();
inspect_rx.await.unwrap();
assert!(
tokio::time::timeout(Duration::from_millis(50), socket.read_u8())
.await
.is_err(),
"oversized rejection must leave the TCP stream empty"
);
inspected_tx.send(()).unwrap();
let mut received = vec![0; exact_limit];
socket.read_exact(&mut received).await.unwrap();
received
});
let transport = TcpTransport::builder()
.max_message_size(4096)
.send_capacity(exact_limit)
.connect(server_addr)
.await
.unwrap();
assert_eq!(transport.receive_limits().advertised().as_usize(), 4096);
assert_eq!(transport.send_capacity(), exact_limit);
let mut oversized = request.clone();
oversized.push(0);
let error = transport.send(&oversized).await.unwrap_err();
assert!(matches!(
*error,
Error::OutboundMessageTooLarge { size, limit }
if size == exact_limit + 1 && limit == exact_limit
));
inspect_tx.send(()).unwrap();
inspected_rx.await.unwrap();
transport.send(&request).await.unwrap();
assert_eq!(server.await.unwrap(), request);
}
#[tokio::test]
async fn test_tcp_is_reliable() {
let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
let server_addr = listener.local_addr().unwrap();
tokio::spawn(async move {
let _ = listener.accept().await;
});
let transport = TcpTransport::connect(server_addr).await.unwrap();
assert!(transport.is_reliable());
}
#[tokio::test]
async fn test_tcp_concurrent_requests() {
let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
let server_addr = listener.local_addr().unwrap();
let server = tokio::spawn(async move {
let (mut socket, _) = listener.accept().await.unwrap();
for _ in 0..5 {
let mut tag = [0u8; 1];
if socket.read_exact(&mut tag).await.is_err() {
break;
}
let mut len_byte = [0u8; 1];
socket.read_exact(&mut len_byte).await.unwrap();
let content_len = len_byte[0] as usize;
let mut content = vec![0u8; content_len];
socket.read_exact(&mut content).await.unwrap();
let mut request = Vec::with_capacity(content_len + 2);
request.extend_from_slice(&tag);
request.extend_from_slice(&len_byte);
request.extend_from_slice(&content);
let request_id = extract_request_id(&request).unwrap();
let response = build_response_with_id(request_id);
socket.write_all(&response).await.unwrap();
}
});
let transport = TcpTransport::connect(server_addr).await.unwrap();
let mut handles = vec![];
for i in 0..5 {
let transport = transport.clone();
let handle = tokio::spawn(async move {
let request_id = i + 1;
let request = build_request_with_id(request_id);
let registration =
RequestRegistration::test_unchecked(request_id, Duration::from_secs(5));
let (response, _) = transport.request(&request, registration).await?;
assert_eq!(response[0], 0x30, "Response should be SEQUENCE");
Ok::<_, Box<Error>>(i)
});
handles.push(handle);
}
let results: Vec<_> = futures::future::join_all(handles).await;
let success_count = results
.iter()
.filter(|r| r.as_ref().is_ok_and(std::result::Result::is_ok))
.count();
assert_eq!(
success_count, 5,
"All 5 concurrent requests should succeed (serialized)"
);
server.await.unwrap();
}
fn build_request_with_id(request_id: i32) -> Vec<u8> {
let id_bytes = request_id.to_be_bytes();
vec![
0x30,
0x1d, 0x02,
0x01,
0x01, 0x04,
0x06,
0x70,
0x75,
0x62,
0x6c,
0x69,
0x63, 0xa0,
0x10, 0x02,
0x04,
id_bytes[0],
id_bytes[1],
id_bytes[2],
id_bytes[3], 0x02,
0x01,
0x00, 0x02,
0x01,
0x00, 0x30,
0x02,
0x30,
0x00, ]
}
fn build_response_with_id(request_id: i32) -> Vec<u8> {
let id_bytes = request_id.to_be_bytes();
vec![
0x30,
0x1d, 0x02,
0x01,
0x01, 0x04,
0x06,
0x70,
0x75,
0x62,
0x6c,
0x69,
0x63, 0xa2,
0x10, 0x02,
0x04,
id_bytes[0],
id_bytes[1],
id_bytes[2],
id_bytes[3], 0x02,
0x01,
0x00, 0x02,
0x01,
0x00, 0x30,
0x02,
0x30,
0x00, ]
}
#[tokio::test]
async fn tcp_skips_mismatched_community_frame_under_original_deadline() {
let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
let addr = listener.local_addr().unwrap();
let server = tokio::spawn(async move {
let (mut stream, _) = listener.accept().await.unwrap();
let mut request = [0u8; 31];
stream.read_exact(&mut request).await.unwrap();
let mut wrong = build_response_with_id(77);
wrong[7..13].copy_from_slice(b"other!");
stream.write_all(&wrong).await.unwrap();
stream.write_all(&build_response_with_id(77)).await.unwrap();
});
let transport = TcpTransport::connect(addr).await.unwrap();
let registration = RequestRegistration::community(
77,
deadline_after(Duration::from_secs(2)),
crate::CommunityVersion::V2c,
Bytes::from_static(b"public"),
super::super::CommunityResponsePolicy::Exact,
);
let (response, _) = transport
.request(&build_request_with_id(77), registration)
.await
.unwrap();
assert_eq!(response.as_ref(), build_response_with_id(77));
assert!(!transport.inner.is_poisoned());
server.await.unwrap();
}
#[tokio::test]
async fn tcp_strict_envelope_policy_applies_to_each_framed_message() {
let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
let addr = listener.local_addr().unwrap();
let server = tokio::spawn(async move {
let (mut stream, _) = listener.accept().await.unwrap();
let mut request = [0u8; 31];
stream.read_exact(&mut request).await.unwrap();
stream.write_all(&build_response_with_id(77)).await.unwrap();
stream.write_all(&build_response_with_id(78)).await.unwrap();
});
let transport = TcpTransport::connect(addr).await.unwrap();
let registration = RequestRegistration::community(
77,
deadline_after(Duration::from_secs(2)),
crate::CommunityVersion::V2c,
Bytes::from_static(b"public"),
super::super::CommunityResponsePolicy::Exact,
)
.with_decode_config(crate::DecodeConfig::STRICT);
let (response, _) = transport
.request(&build_request_with_id(77), registration)
.await
.unwrap();
assert_eq!(response.as_ref(), build_response_with_id(77));
assert!(!transport.inner.is_poisoned());
server.await.unwrap();
}
#[tokio::test]
async fn tcp_skips_stale_id_then_accepts_primary_id() {
let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
let addr = listener.local_addr().unwrap();
let server = tokio::spawn(async move {
let (mut stream, _) = listener.accept().await.unwrap();
let mut request = [0u8; 31];
stream.read_exact(&mut request).await.unwrap();
stream.write_all(&build_response_with_id(90)).await.unwrap();
stream.write_all(&build_response_with_id(91)).await.unwrap();
});
let transport = TcpTransport::connect(addr).await.unwrap();
let (response, _) = transport
.request(
&build_request_with_id(91),
RequestRegistration::test_unchecked(91, Duration::from_secs(2)),
)
.await
.unwrap();
assert_eq!(extract_request_id(&response), Some(91));
assert!(!transport.inner.is_poisoned());
server.await.unwrap();
}
#[tokio::test]
async fn tcp_accepts_each_registered_alias_id() {
for accepted_alias in [101, 102] {
let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
let addr = listener.local_addr().unwrap();
let server = tokio::spawn(async move {
let (mut stream, _) = listener.accept().await.unwrap();
let mut request = [0u8; 31];
stream.read_exact(&mut request).await.unwrap();
stream
.write_all(&build_response_with_id(accepted_alias))
.await
.unwrap();
});
let transport = TcpTransport::connect(addr).await.unwrap();
let registration = RequestRegistration::test_unchecked(100, Duration::from_secs(2))
.with_aliases([101, 102])
.unwrap();
let (response, _) = transport
.request(&build_request_with_id(100), registration)
.await
.unwrap();
assert_eq!(extract_request_id(&response), Some(accepted_alias));
assert!(!transport.inner.is_poisoned());
server.await.unwrap();
}
}
#[tokio::test]
async fn stale_frames_do_not_extend_deadline_and_timeout_poisons() {
let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
let addr = listener.local_addr().unwrap();
let server = tokio::spawn(async move {
let (mut stream, _) = listener.accept().await.unwrap();
let mut request = [0u8; 31];
stream.read_exact(&mut request).await.unwrap();
for stale_id in 200..230 {
if stream
.write_all(&build_response_with_id(stale_id))
.await
.is_err()
{
break;
}
tokio::time::sleep(Duration::from_millis(10)).await;
}
});
let transport = TcpTransport::connect(addr).await.unwrap();
let start = tokio::time::Instant::now();
let error = transport
.request(
&build_request_with_id(300),
RequestRegistration::v3(300, deadline_after(Duration::from_millis(100))),
)
.await
.unwrap_err();
let elapsed = start.elapsed();
assert!(matches!(*error, Error::Timeout { .. }));
assert!(elapsed < Duration::from_millis(250), "elapsed {elapsed:?}");
assert!(transport.inner.is_poisoned());
server.abort();
}
#[tokio::test]
async fn tcp_validator_rejection_keeps_exchange_for_later_frame() {
let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
let addr = listener.local_addr().unwrap();
let server = tokio::spawn(async move {
let (mut stream, _) = listener.accept().await.unwrap();
let mut request = [0u8; 31];
stream.read_exact(&mut request).await.unwrap();
stream.write_all(&build_response_with_id(78)).await.unwrap();
stream.write_all(&build_response_with_id(78)).await.unwrap();
});
let transport = TcpTransport::connect(addr).await.unwrap();
let registration = RequestRegistration::community(
78,
deadline_after(Duration::from_secs(2)),
crate::CommunityVersion::V2c,
Bytes::from_static(b"public"),
super::super::CommunityResponsePolicy::Exact,
);
let mut candidates = 0;
let response = transport
.request_with(&build_request_with_id(78), registration, |data, _| {
candidates += 1;
if candidates == 1 {
Ok(Candidate::Reject)
} else {
Ok(Candidate::Accept(data))
}
})
.await
.unwrap();
assert_eq!(response.as_ref(), build_response_with_id(78));
assert_eq!(candidates, 2);
assert!(!transport.inner.is_poisoned());
server.await.unwrap();
}
#[tokio::test]
async fn test_tcp_rejects_excessive_claimed_size() {
let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
let server_addr = listener.local_addr().unwrap();
let server = tokio::spawn(async move {
let (mut socket, _) = listener.accept().await.unwrap();
let mut buf = [0u8; 64];
let _ = socket.read(&mut buf).await;
let malicious_response = [
0x30, 0x84, 0x06, 0x40, 0x00,
0x00, ];
let _ = socket.write_all(&malicious_response).await;
tokio::time::sleep(Duration::from_millis(100)).await;
});
let transport = TcpTransport::connect(server_addr).await.unwrap();
let request = build_request_with_id(1);
transport.send(&request).await.unwrap();
let registration = RequestRegistration::v3(1, deadline_after(Duration::from_secs(5)));
let result = transport.recv(registration).await;
assert!(result.is_err(), "Should reject excessive claimed size");
let err = result.unwrap_err();
assert!(
matches!(*err, Error::Decode(_)),
"Expected Decode error, got: {err:?}"
);
server.await.unwrap();
}
#[tokio::test]
async fn test_read_ber_message_rejects_bad_tag() {
let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
let server_addr = listener.local_addr().unwrap();
let server = tokio::spawn(async move {
let (mut socket, _) = listener.accept().await.unwrap();
socket.write_all(&[0x31, 0x00]).await.unwrap();
});
let mut client = TcpStream::connect(server_addr).await.unwrap();
let result = timeout(
Duration::from_secs(5),
read_ber_message(&mut client, server_addr, DEFAULT_MAX_MESSAGE_SIZE),
)
.await
.expect("read_ber_message should not hang");
assert!(result.is_err(), "Should reject non-0x30 tag byte");
let err = result.unwrap_err();
assert!(
matches!(&*err, Error::Decode(error)
if error.offset == 0
&& error.kind == crate::DecodeErrorKind::UnexpectedTag { expected: 0x30, actual: 0x31 }
&& error.peer == Some(server_addr)),
"Expected Decode error, got: {err:?}"
);
server.await.unwrap();
}
#[tokio::test]
async fn tcp_receive_reports_high_tag_number_form_at_network_boundary() {
let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
let server_addr = listener.local_addr().unwrap();
let server = tokio::spawn(async move {
let (mut socket, _) = listener.accept().await.unwrap();
socket.write_all(&[0xbf]).await.unwrap();
});
let transport = TcpTransport::connect(server_addr).await.unwrap();
let error = transport
.recv(RequestRegistration::v3(
1,
deadline_after(Duration::from_secs(5)),
))
.await
.expect_err("high-tag-number form must be rejected");
assert!(
matches!(&*error, Error::Decode(error)
if error.kind == crate::DecodeErrorKind::UnsupportedMultiOctetTag { first_octet: 0xbf }
&& error.origin == crate::DecodeErrorOrigin::Packet
&& error.offset == 0
&& error.peer == Some(server_addr)),
"unexpected error: {error:?}"
);
server.await.unwrap();
}
#[tokio::test]
async fn test_read_ber_message_rejects_indefinite_length() {
let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
let server_addr = listener.local_addr().unwrap();
let server = tokio::spawn(async move {
let (mut socket, _) = listener.accept().await.unwrap();
socket.write_all(&[0x30, 0x80]).await.unwrap();
});
let mut client = TcpStream::connect(server_addr).await.unwrap();
let result = timeout(
Duration::from_secs(5),
read_ber_message(&mut client, server_addr, DEFAULT_MAX_MESSAGE_SIZE),
)
.await
.expect("read_ber_message should not hang");
assert!(result.is_err(), "Should reject indefinite length encoding");
let err = result.unwrap_err();
assert!(
matches!(*err, Error::Decode(_)),
"Expected Decode error, got: {err:?}"
);
server.await.unwrap();
}
#[tokio::test]
async fn test_read_ber_message_rejects_length_encoding_too_long() {
let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
let server_addr = listener.local_addr().unwrap();
let server = tokio::spawn(async move {
let (mut socket, _) = listener.accept().await.unwrap();
socket
.write_all(&[0x30, 0x85, 0x00, 0x00, 0x00, 0x00, 0x00])
.await
.unwrap();
});
let mut client = TcpStream::connect(server_addr).await.unwrap();
let result = timeout(
Duration::from_secs(5),
read_ber_message(&mut client, server_addr, DEFAULT_MAX_MESSAGE_SIZE),
)
.await
.expect("read_ber_message should not hang");
assert!(
result.is_err(),
"Should reject length encoding over 4 octets"
);
let err = result.unwrap_err();
assert!(
matches!(*err, Error::Decode(_)),
"Expected Decode error, got: {err:?}"
);
server.await.unwrap();
}
#[tokio::test]
async fn test_read_ber_message_reassembles_segmented_content() {
let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
let server_addr = listener.local_addr().unwrap();
let server = tokio::spawn(async move {
let (mut socket, _) = listener.accept().await.unwrap();
socket.write_all(&[0x30, 0x04]).await.unwrap();
socket.flush().await.unwrap();
tokio::time::sleep(Duration::from_millis(20)).await;
socket.write_all(&[0x01, 0x02]).await.unwrap();
socket.flush().await.unwrap();
tokio::time::sleep(Duration::from_millis(20)).await;
socket.write_all(&[0x03, 0x04]).await.unwrap();
socket.flush().await.unwrap();
tokio::time::sleep(Duration::from_millis(100)).await;
});
let mut client = TcpStream::connect(server_addr).await.unwrap();
let result = timeout(
Duration::from_secs(5),
read_ber_message(&mut client, server_addr, DEFAULT_MAX_MESSAGE_SIZE),
)
.await
.expect("read_ber_message should not hang");
let bytes = result.expect("segmented message should reassemble successfully");
assert_eq!(bytes.as_ref(), &[0x30, 0x04, 0x01, 0x02, 0x03, 0x04]);
server.await.unwrap();
}
#[tokio::test]
async fn test_read_ber_message_truncated_stream_is_network_error() {
let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
let server_addr = listener.local_addr().unwrap();
let server = tokio::spawn(async move {
let (mut socket, _) = listener.accept().await.unwrap();
socket.write_all(&[0x30, 0x05]).await.unwrap();
socket.write_all(&[0x01, 0x02]).await.unwrap();
socket.flush().await.unwrap();
drop(socket);
});
let mut client = TcpStream::connect(server_addr).await.unwrap();
let result = timeout(
Duration::from_secs(5),
read_ber_message(&mut client, server_addr, DEFAULT_MAX_MESSAGE_SIZE),
)
.await
.expect("read_ber_message should not hang");
assert!(result.is_err(), "Should error on truncated content stream");
let err = result.unwrap_err();
assert!(
matches!(*err, Error::Network { .. }),
"Expected Network error, got: {err:?}"
);
server.await.unwrap();
}
#[tokio::test]
async fn invalid_tcp_message_sizes_are_rejected_before_connect() {
let target = "127.0.0.1:9".parse().unwrap();
for size in [0usize, 483, i32::MAX as usize + 1, usize::MAX] {
let error = TcpTransport::builder()
.max_message_size(size)
.connect(target)
.await
.err()
.expect("invalid size must fail before connect");
assert!(
matches!(*error, Error::Config(_)),
"unexpected error: {error}"
);
}
}
#[tokio::test]
async fn tcp_frame_limit_counts_total_encoded_size() {
const CONTENT_LEN: usize = 481;
const TOTAL_LEN: usize = 1 + 3 + CONTENT_LEN;
let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
let server_addr = listener.local_addr().unwrap();
let server = tokio::spawn(async move {
let (mut socket, _) = listener.accept().await.unwrap();
let mut frame = vec![0x30, 0x82, 0x01, 0xe1];
frame.resize(TOTAL_LEN, 0);
socket.write_all(&frame).await.unwrap();
});
let mut client = TcpStream::connect(server_addr).await.unwrap();
let frame = read_ber_message(&mut client, server_addr, TOTAL_LEN)
.await
.expect("exact total-message limit should be accepted");
assert_eq!(frame.len(), TOTAL_LEN);
server.await.unwrap();
let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
let server_addr = listener.local_addr().unwrap();
let server = tokio::spawn(async move {
let (mut socket, _) = listener.accept().await.unwrap();
socket.write_all(&[0x30, 0x82, 0x01, 0xe1]).await.unwrap();
});
let mut client = TcpStream::connect(server_addr).await.unwrap();
let error = read_ber_message(&mut client, server_addr, TOTAL_LEN - 1)
.await
.expect_err("one byte over total-message limit must be rejected");
assert!(matches!(*error, Error::Decode(_)));
server.await.unwrap();
}
#[tokio::test]
async fn test_tcp_builder_custom_message_limit() {
let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
let server_addr = listener.local_addr().unwrap();
let server = tokio::spawn(async move {
let (mut socket, _) = listener.accept().await.unwrap();
let mut buf = [0u8; 64];
let _ = socket.read(&mut buf).await;
let response = [
0x30, 0x82, 0x28, 0x00, ];
let _ = socket.write_all(&response).await;
tokio::time::sleep(Duration::from_millis(100)).await;
});
let transport = TcpTransport::builder()
.max_message_size(1024) .connect(server_addr)
.await
.unwrap();
let request = build_request_with_id(1);
transport.send(&request).await.unwrap();
let registration = RequestRegistration::v3(1, deadline_after(Duration::from_secs(5)));
let result = transport.recv(registration).await;
assert!(
result.is_err(),
"Should reject message exceeding custom limit"
);
let err = result.unwrap_err();
assert!(
matches!(*err, Error::Decode(_)),
"Expected Decode error, got: {err:?}"
);
server.await.unwrap();
}
#[tokio::test]
async fn cancellation_before_stream_lock_does_not_poison() {
let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
let server_addr = listener.local_addr().unwrap();
let server = tokio::spawn(async move {
let (_socket, _) = listener.accept().await.unwrap();
tokio::time::sleep(Duration::from_secs(1)).await;
});
let transport = TcpTransport::connect(server_addr).await.unwrap();
let held_lock = transport.inner.stream.clone().lock_owned().await;
let request = build_request_with_id(1);
let registration = RequestRegistration::v3(1, deadline_after(Duration::from_secs(5)));
let cancelled = timeout(
Duration::from_millis(30),
transport.request(&request, registration),
)
.await;
assert!(cancelled.is_err());
assert!(!transport.inner.is_poisoned());
drop(held_lock);
server.abort();
}
#[tokio::test(start_paused = true)]
async fn queued_request_deadline_includes_lock_wait_without_poisoning() {
let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
let server_addr = listener.local_addr().unwrap();
let server = tokio::spawn(async move {
let (_socket, _) = listener.accept().await.unwrap();
std::future::pending::<()>().await;
});
let transport = TcpTransport::connect(server_addr).await.unwrap();
let held_lock = transport.inner.stream.clone().lock_owned().await;
assert!(!transport.inner.is_poisoned());
let second_data = build_request_with_id(2);
let mut second = Box::pin(transport.request(
&second_data,
RequestRegistration::v3(2, deadline_after(Duration::from_secs(5))),
));
assert!(matches!(futures::poll!(second.as_mut()), Poll::Pending));
assert!(!transport.inner.is_poisoned());
tokio::time::advance(Duration::from_secs(5)).await;
let error = second.await.unwrap_err();
assert!(matches!(
*error,
Error::Timeout {
elapsed,
retries: 0,
..
} if elapsed == Duration::from_secs(5)
));
assert!(
!transport.inner.is_poisoned(),
"expiry while queued did not touch the stream"
);
drop(held_lock);
server.abort();
}
#[tokio::test(start_paused = true)]
async fn request_uses_remaining_budget_after_lock_wait() {
let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
let server_addr = listener.local_addr().unwrap();
let (request_read_tx, request_read_rx) = tokio::sync::oneshot::channel();
let server = tokio::spawn(async move {
let (mut socket, _) = listener.accept().await.unwrap();
let mut request = [0u8; 31];
socket.read_exact(&mut request).await.unwrap();
request_read_tx.send(()).unwrap();
std::future::pending::<()>().await;
});
let transport = TcpTransport::connect(server_addr).await.unwrap();
let held_lock = transport.inner.stream.clone().lock_owned().await;
let request_data = build_request_with_id(1);
let mut request = Box::pin(transport.request(
&request_data,
RequestRegistration::v3(1, deadline_after(Duration::from_secs(10))),
));
assert!(matches!(futures::poll!(request.as_mut()), Poll::Pending));
tokio::time::advance(Duration::from_secs(6)).await;
drop(held_lock);
assert!(matches!(futures::poll!(request.as_mut()), Poll::Pending));
request_read_rx.await.unwrap();
tokio::time::advance(Duration::from_secs(3)).await;
assert!(matches!(futures::poll!(request.as_mut()), Poll::Pending));
tokio::time::advance(Duration::from_secs(1)).await;
let error = request.await.unwrap_err();
assert!(matches!(*error, Error::Timeout { .. }));
assert!(transport.inner.is_poisoned());
server.abort();
}
#[tokio::test(start_paused = true)]
async fn exact_deadline_boundary_after_lock_acquisition_does_not_poison() {
let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
let server_addr = listener.local_addr().unwrap();
let server = tokio::spawn(async move {
let (mut socket, _) = listener.accept().await.unwrap();
assert!(
tokio::time::timeout(Duration::from_secs(1), socket.read_u8())
.await
.is_err(),
"an exact-boundary timeout must not write"
);
});
let transport = TcpTransport::connect(server_addr).await.unwrap();
let held_lock = transport.inner.stream.clone().lock_owned().await;
let data = build_request_with_id(1);
let mut request = Box::pin(transport.request(
&data,
RequestRegistration::v3(1, deadline_after(Duration::from_secs(5))),
));
assert!(matches!(futures::poll!(request.as_mut()), Poll::Pending));
tokio::time::advance(Duration::from_secs(5)).await;
drop(held_lock);
let error = request.await.unwrap_err();
assert!(matches!(*error, Error::Timeout { .. }));
assert!(
!transport.inner.is_poisoned(),
"deadline won before the stream I/O future was polled"
);
server.await.unwrap();
}
#[tokio::test(start_paused = true)]
async fn standalone_send_lock_timeout_and_cancellation_do_not_poison() {
let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
let server_addr = listener.local_addr().unwrap();
let server = tokio::spawn(async move {
let (mut socket, _) = listener.accept().await.unwrap();
let mut request = [0u8; 31];
socket.read_exact(&mut request).await.unwrap();
});
let transport = TcpTransport::connect(server_addr).await.unwrap();
let held_lock = transport.inner.stream.clone().lock_owned().await;
let cancelled_transport = transport.clone();
let cancelled = tokio::spawn(async move {
cancelled_transport
.send_with_timeout(&build_request_with_id(1), Duration::from_secs(30))
.await
});
tokio::task::yield_now().await;
cancelled.abort();
assert!(cancelled.await.unwrap_err().is_cancelled());
assert!(!transport.inner.is_poisoned());
let timed_data = build_request_with_id(2);
let mut timed = Box::pin(transport.send_with_timeout(&timed_data, Duration::from_secs(5)));
assert!(matches!(futures::poll!(timed.as_mut()), Poll::Pending));
tokio::time::advance(Duration::from_secs(5)).await;
let error = timed.await.unwrap_err();
assert!(matches!(*error, Error::Timeout { .. }));
assert!(!transport.inner.is_poisoned());
drop(held_lock);
transport
.send_with_timeout(&build_request_with_id(3), Duration::from_secs(5))
.await
.unwrap();
server.await.unwrap();
}
#[cfg(target_os = "linux")]
#[tokio::test(start_paused = true)]
async fn standalone_send_write_timeout_poisons_after_partial_io() {
let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
let server_addr = listener.local_addr().unwrap();
let (touched_tx, touched_rx) = tokio::sync::oneshot::channel();
let server = tokio::spawn(async move {
let (mut socket, _) = listener.accept().await.unwrap();
socket.read_u8().await.unwrap();
touched_tx.send(()).unwrap();
std::future::pending::<()>().await;
});
let socket = tokio::net::TcpSocket::new_v4().unwrap();
socket.set_send_buffer_size(1024).unwrap();
let transport = TcpTransport::from_socket(
socket,
server_addr,
TcpOptions {
send_capacity: 16 * 1024 * 1024,
..TcpOptions::default()
},
)
.await
.unwrap();
let send_transport = transport.clone();
let send = tokio::spawn(async move {
let data = vec![0xaa; 16 * 1024 * 1024];
send_transport
.send_with_timeout(&data, Duration::from_secs(10))
.await
});
touched_rx.await.unwrap();
tokio::task::yield_now().await;
assert!(!send.is_finished(), "large write must remain partial");
tokio::time::advance(Duration::from_secs(10)).await;
let error = send.await.unwrap().unwrap_err();
assert!(matches!(*error, Error::Timeout { .. }));
assert!(transport.inner.is_poisoned());
server.abort();
}
#[tokio::test]
async fn cancellation_after_complete_request_write_poisons() {
let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
let server_addr = listener.local_addr().unwrap();
let (written_tx, written_rx) = tokio::sync::oneshot::channel();
let server = tokio::spawn(async move {
let (mut socket, _) = listener.accept().await.unwrap();
let mut request = [0u8; 31];
socket.read_exact(&mut request).await.unwrap();
written_tx.send(()).unwrap();
tokio::time::sleep(Duration::from_secs(1)).await;
});
let transport = TcpTransport::connect(server_addr).await.unwrap();
let request_transport = transport.clone();
let request = tokio::spawn(async move {
request_transport
.request(
&build_request_with_id(1),
RequestRegistration::v3(1, deadline_after(Duration::from_secs(30))),
)
.await
});
written_rx.await.unwrap();
request.abort();
assert!(request.await.unwrap_err().is_cancelled());
assert!(transport.inner.is_poisoned());
let error = transport
.request(
&build_request_with_id(2),
RequestRegistration::v3(2, deadline_after(Duration::from_secs(1))),
)
.await
.unwrap_err();
assert!(matches!(*error, Error::Closed { .. }));
server.abort();
}
#[tokio::test]
async fn cancellation_during_partial_frame_read_poisons() {
let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
let server_addr = listener.local_addr().unwrap();
let (partial_tx, partial_rx) = tokio::sync::oneshot::channel();
let server = tokio::spawn(async move {
let (mut socket, _) = listener.accept().await.unwrap();
let mut request = [0u8; 31];
socket.read_exact(&mut request).await.unwrap();
socket.write_all(&[0x30, 0x08, 0x02]).await.unwrap();
partial_tx.send(()).unwrap();
tokio::time::sleep(Duration::from_secs(1)).await;
});
let transport = TcpTransport::connect(server_addr).await.unwrap();
let request_transport = transport.clone();
let request = tokio::spawn(async move {
request_transport
.request(
&build_request_with_id(1),
RequestRegistration::v3(1, deadline_after(Duration::from_secs(30))),
)
.await
});
partial_rx.await.unwrap();
tokio::task::yield_now().await;
request.abort();
assert!(request.await.unwrap_err().is_cancelled());
assert!(transport.inner.is_poisoned());
server.abort();
}
#[cfg(target_os = "linux")]
#[tokio::test]
async fn cancellation_during_partial_send_poisons() {
let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
let server_addr = listener.local_addr().unwrap();
let (accepted_tx, accepted_rx) = tokio::sync::oneshot::channel();
let server = tokio::spawn(async move {
let (_socket, _) = listener.accept().await.unwrap();
accepted_tx.send(()).unwrap();
tokio::time::sleep(Duration::from_secs(2)).await;
});
let socket = tokio::net::TcpSocket::new_v4().unwrap();
socket.set_send_buffer_size(1024).unwrap();
let transport = TcpTransport::from_socket(
socket,
server_addr,
TcpOptions {
send_capacity: 16 * 1024 * 1024,
..TcpOptions::default()
},
)
.await
.unwrap();
accepted_rx.await.unwrap();
let send_transport = transport.clone();
let send = tokio::spawn(async move {
let data = vec![0xaa; 16 * 1024 * 1024];
send_transport.send(&data).await
});
tokio::time::sleep(Duration::from_millis(30)).await;
assert!(!send.is_finished(), "test write must still be partial");
send.abort();
assert!(send.await.unwrap_err().is_cancelled());
assert!(transport.inner.is_poisoned());
server.abort();
}
#[tokio::test]
async fn request_deadline_covers_partial_write() {
let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
let server_addr = listener.local_addr().unwrap();
let server = tokio::spawn(async move {
let (_socket, _) = listener.accept().await.unwrap();
tokio::time::sleep(Duration::from_secs(2)).await;
});
let socket = tokio::net::TcpSocket::new_v4().unwrap();
socket.set_send_buffer_size(1024).unwrap();
let transport = TcpTransport::from_socket(
socket,
server_addr,
TcpOptions {
send_capacity: 16 * 1024 * 1024,
..TcpOptions::default()
},
)
.await
.unwrap();
let data = vec![0xaa; 16 * 1024 * 1024];
let start = tokio::time::Instant::now();
let error = transport
.request(
&data,
RequestRegistration::v3(1, deadline_after(Duration::from_millis(50))),
)
.await
.unwrap_err();
assert!(matches!(*error, Error::Timeout { .. }));
assert!(start.elapsed() < Duration::from_millis(500));
assert!(transport.inner.is_poisoned());
server.abort();
}
#[tokio::test]
async fn test_tcp_malformed_frame_poisons_stream() {
let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
let server_addr = listener.local_addr().unwrap();
let server = tokio::spawn(async move {
let (mut socket, _) = listener.accept().await.unwrap();
let mut hdr = [0u8; 2];
socket.read_exact(&mut hdr).await.unwrap();
let mut body = vec![0u8; hdr[1] as usize];
socket.read_exact(&mut body).await.unwrap();
socket
.write_all(&[0x31, 0x02, 0xde, 0xad, 0x30, 0x00])
.await
.unwrap();
socket.flush().await.unwrap();
tokio::time::sleep(Duration::from_millis(200)).await;
});
let transport = TcpTransport::connect(server_addr).await.unwrap();
let request = build_request_with_id(1);
let registration = RequestRegistration::v3(1, deadline_after(Duration::from_secs(5)));
let first = transport.request(&request, registration).await;
let err = first.expect_err("malformed frame should error");
assert!(
matches!(*err, Error::Decode(_)),
"Expected Decode error, got: {err:?}"
);
assert!(
transport.inner.is_poisoned(),
"stream should be poisoned after a malformed frame"
);
let request2 = build_request_with_id(2);
let registration = RequestRegistration::v3(2, deadline_after(Duration::from_secs(5)));
let second = timeout(
Duration::from_secs(5),
transport.request(&request2, registration),
)
.await
.expect("second request should not hang");
let err2 = second.expect_err("poisoned stream should reject the next request");
assert!(
matches!(*err2, Error::Closed { .. }),
"Expected Closed on poisoned stream, got: {err2:?}"
);
server.await.unwrap();
}
#[tokio::test]
async fn test_tcp_receive_uses_owned_registration_timeout() {
let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
let server_addr = listener.local_addr().unwrap();
let server = tokio::spawn(async move {
let (_socket, _) = listener.accept().await.unwrap();
tokio::time::sleep(Duration::from_secs(30)).await;
});
let transport = TcpTransport::connect(server_addr).await.unwrap();
let registration = RequestRegistration::v3(1, deadline_after(Duration::from_millis(150)));
let start = std::time::Instant::now();
let result = timeout(Duration::from_secs(5), transport.recv(registration))
.await
.expect("recv should honor the owned short timeout");
let elapsed = start.elapsed();
let err = result.expect_err("recv should time out");
assert!(
matches!(*err, Error::Timeout { .. }),
"Expected Timeout, got: {err:?}"
);
assert!(
elapsed < Duration::from_secs(2),
"recv honored the wrong timeout; elapsed {elapsed:?}"
);
server.abort();
}
#[tokio::test]
async fn test_tcp_timeout_mid_frame_poisons_stream() {
let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
let server_addr = listener.local_addr().unwrap();
let server = tokio::spawn(async move {
let (mut socket, _) = listener.accept().await.unwrap();
let mut buf = [0u8; 64];
let _ = socket.read(&mut buf).await;
socket.write_all(&[0x30, 0x08, 0x01, 0x02]).await.unwrap();
socket.flush().await.unwrap();
tokio::time::sleep(Duration::from_millis(500)).await;
});
let transport = TcpTransport::connect(server_addr).await.unwrap();
let request = build_request_with_id(1);
transport.send(&request).await.unwrap();
let registration = RequestRegistration::v3(1, deadline_after(Duration::from_millis(100)));
let first = transport.recv(registration).await;
let err = first.expect_err("mid-frame read should time out");
assert!(
matches!(*err, Error::Timeout { .. }),
"Expected Timeout, got: {err:?}"
);
assert!(
transport.inner.is_poisoned(),
"stream should be poisoned after a mid-frame timeout"
);
let registration = RequestRegistration::v3(1, deadline_after(Duration::from_secs(5)));
let second = timeout(Duration::from_secs(5), transport.recv(registration))
.await
.expect("second recv should not hang");
let err2 = second.expect_err("poisoned stream should reject the next recv");
assert!(
matches!(*err2, Error::Closed { .. }),
"Expected Closed on poisoned stream, got: {err2:?}"
);
server.await.unwrap();
}
}