pub mod auth;
pub mod header;
use crate::error::{NfsError, Result};
use byteorder::{BigEndian, ByteOrder};
use bytes::{Bytes, BytesMut};
use std::collections::HashMap;
use std::future::Future;
use std::net::SocketAddr;
use std::pin::Pin;
use std::sync::Arc;
use std::sync::atomic::{AtomicBool, AtomicU8, AtomicU32, AtomicU64, Ordering};
use tokio::io::{AsyncReadExt, AsyncWriteExt, BufReader};
use tokio::net::tcp::{OwnedReadHalf, OwnedWriteHalf};
use tokio::sync::{Mutex as TokioMutex, Notify, oneshot};
use tokio::task::JoinHandle;
use tracing::{debug, error, info, trace, warn};
use auth::Auth;
pub(crate) use header::Header;
pub(crate) const RPC_VERSION: u32 = 2;
pub(crate) const PORTMAP_PROG: u32 = 100000;
pub(crate) const PORTMAP_VERSION: u32 = 2;
pub(crate) const PORTMAP_PORT: u16 = 111;
pub(crate) const MOUNT_PROG: u32 = 100005;
pub(crate) const MOUNT3_VERSION: u32 = 3;
pub(crate) const NFS_PROG: u32 = 100003;
pub(crate) const NFS3_VERSION: u32 = 3;
const IPPROTO_TCP: u32 = 6;
const METADATA_TIMEOUT: std::time::Duration = std::time::Duration::from_secs(5);
#[derive(Clone, Copy, Debug, Eq, PartialEq)]
pub(crate) struct ReplayPolicy {
max_attempts: usize,
}
impl ReplayPolicy {
pub(crate) const ONE_ATTEMPT: Self = Self { max_attempts: 1 };
pub(crate) const fn byte_identical(max_attempts: usize) -> Self {
assert!(
max_attempts > 1,
"byte-identical RPC replay requires at least 2 attempts"
);
Self { max_attempts }
}
const fn max_attempts(self) -> usize {
self.max_attempts
}
}
enum PortmapProc2 {
Null = 0,
GetPort = 3,
}
pub(crate) async fn portmap(
addrs: &Vec<SocketAddr>,
prog: u32,
vers: u32,
auth: &Auth,
max_retries: usize,
noresvport: bool,
) -> Result<u16> {
let mut last_err: Option<(SocketAddr, NfsError)> = None;
for addr in addrs {
debug!(addr = %addr, prog, vers, "attempting portmapper lookup");
match portmap_on_addr(addr, prog, vers, auth, max_retries, noresvport).await {
Ok(port) => {
info!(addr = %addr, prog, vers, port, "portmapper resolved port");
return Ok(port);
}
Err(e) => {
warn!(addr = %addr, prog, vers, error = %e, "portmapper lookup failed on address");
last_err = Some((*addr, e));
}
}
}
Err(NfsError::Rpc(format!(
"portmapper lookup failed for prog={} vers={}: {}",
prog,
vers,
last_err
.map(|(addr, e)| format!("{}: {}", addr, e))
.unwrap_or_else(|| "no addresses tried".to_string()),
)))
}
async fn portmap_on_addr(
addr: &SocketAddr,
prog: u32,
vers: u32,
auth: &Auth,
max_retries: usize,
noresvport: bool,
) -> Result<u16> {
let mux = StreamMux::connect(*addr, noresvport).await?;
let client = Client::new(mux, None);
let result = portmap_calls(&client, prog, vers, auth, max_retries).await;
let _ = client.shutdown().await;
result
}
async fn portmap_calls(
client: &Client,
prog: u32,
vers: u32,
auth: &Auth,
max_retries: usize,
) -> Result<u16> {
let mut buf = Vec::<u8>::new();
Header::new(
RPC_VERSION,
PORTMAP_PROG,
PORTMAP_VERSION,
PortmapProc2::Null as u32,
auth,
&Auth::new_null(),
)
.encode(&mut buf);
client
.call(
buf,
ReplayPolicy::byte_identical(max_retries),
METADATA_TIMEOUT,
)
.await?;
let args = GETPORT2args {
header: Header::new(
RPC_VERSION,
PORTMAP_PROG,
PORTMAP_VERSION,
PortmapProc2::GetPort as u32,
auth,
&Auth::new_null(),
),
prog,
vers,
proto: IPPROTO_TCP,
port: 0,
};
let mut buf = Vec::<u8>::new();
args.encode(&mut buf);
let res = client
.call(
buf,
ReplayPolicy::byte_identical(max_retries),
METADATA_TIMEOUT,
)
.await?;
Ok(BigEndian::read_u32(&res[..4]) as u16)
}
#[derive(Debug, PartialEq)]
struct GETPORT2args {
header: Header,
prog: u32,
vers: u32,
proto: u32,
port: u32,
}
impl GETPORT2args {
fn encode(&self, buf: &mut Vec<u8>) {
self.header.encode(buf);
buf.extend_from_slice(&self.prog.to_be_bytes());
buf.extend_from_slice(&self.vers.to_be_bytes());
buf.extend_from_slice(&self.proto.to_be_bytes());
buf.extend_from_slice(&self.port.to_be_bytes());
}
}
type PendingMap = Arc<std::sync::Mutex<HashMap<u32, oneshot::Sender<Result<Bytes>>>>>;
struct PendingRequestGuard {
pending: PendingMap,
xid: u32,
}
impl Drop for PendingRequestGuard {
fn drop(&mut self) {
if let Ok(mut map) = self.pending.lock() {
map.remove(&self.xid);
}
}
}
pub(crate) type BackchannelHandler = Arc<dyn Fn(Bytes) -> Option<Vec<u8>> + Send + Sync>;
type ReconnectFuture = Pin<Box<dyn Future<Output = Result<()>> + Send>>;
type ReconnectHandler = Arc<dyn Fn(Client, u64) -> ReconnectFuture + Send + Sync>;
const CONNECTION_READY: u8 = 0;
const CONNECTION_REBINDING: u8 = 1;
const CONNECTION_FAILED: u8 = 2;
struct RebindPublicationGuard<'a> {
readiness: &'a AtomicU8,
notify: &'a Notify,
published: bool,
}
impl RebindPublicationGuard<'_> {
fn publish(mut self) {
self.readiness.store(CONNECTION_READY, Ordering::Release);
self.notify.notify_waiters();
self.published = true;
}
}
impl Drop for RebindPublicationGuard<'_> {
fn drop(&mut self) {
if !self.published {
self.readiness.store(CONNECTION_FAILED, Ordering::Release);
self.notify.notify_waiters();
}
}
}
type BackchannelSlot = Arc<std::sync::Mutex<Option<BackchannelHandler>>>;
pub(crate) struct StreamMux {
writer: Arc<TokioMutex<OwnedWriteHalf>>,
pending: PendingMap,
backchannel: BackchannelSlot,
addr: SocketAddr,
noresvport: bool,
generation: AtomicU64,
reconnect_lock: TokioMutex<()>,
reconnect_handler: std::sync::Mutex<Option<ReconnectHandler>>,
readiness: AtomicU8,
readiness_notify: Notify,
reader_handle: std::sync::Mutex<Option<JoinHandle<()>>>,
shutdown_flag: AtomicBool,
}
impl StreamMux {
pub(crate) async fn connect(addr: SocketAddr, noresvport: bool) -> Result<Arc<Self>> {
let stream = crate::connect_to_target(&addr, noresvport).await?;
let (reader, writer) = stream.into_split();
let pending: PendingMap = Arc::new(std::sync::Mutex::new(HashMap::new()));
let writer = Arc::new(TokioMutex::new(writer));
let backchannel: BackchannelSlot = Arc::new(std::sync::Mutex::new(None));
let reader = BufReader::with_capacity(1_048_576, reader);
let reader_handle = tokio::spawn(reader_loop(
reader,
Arc::clone(&pending),
Arc::clone(&writer),
Arc::clone(&backchannel),
));
info!(addr = %addr, "RPC stream mux connected");
Ok(Arc::new(Self {
writer,
pending,
backchannel,
addr,
noresvport,
generation: AtomicU64::new(0),
reconnect_lock: TokioMutex::new(()),
reconnect_handler: std::sync::Mutex::new(None),
readiness: AtomicU8::new(CONNECTION_READY),
readiness_notify: Notify::new(),
reader_handle: std::sync::Mutex::new(Some(reader_handle)),
shutdown_flag: AtomicBool::new(false),
}))
}
fn generation(&self) -> u64 {
self.generation.load(Ordering::Acquire)
}
fn enable_backchannel(&self, handler: BackchannelHandler) {
match self.backchannel.lock() {
Ok(mut slot) => *slot = Some(handler),
Err(_) => warn!("backchannel slot lock poisoned; cannot enable backchannel"),
}
}
fn set_reconnect_handler(&self, handler: ReconnectHandler) -> Result<()> {
let mut slot = self
.reconnect_handler
.lock()
.map_err(|_| NfsError::Rpc("reconnect handler lock poisoned".to_string()))?;
*slot = Some(handler);
Ok(())
}
async fn wait_until_ready(&self) -> Result<()> {
loop {
match self.readiness.load(Ordering::Acquire) {
CONNECTION_READY => return Ok(()),
CONNECTION_FAILED => {
return Err(NfsError::Rpc(
"NFSv4.1 connection rebind failed; connection is not ready".to_string(),
));
}
_ => {
let notified = self.readiness_notify.notified();
if self.readiness.load(Ordering::Acquire) == CONNECTION_REBINDING {
notified.await;
}
}
}
}
}
async fn send_and_receive_inner(
&self,
xid: u32,
header: &[u8],
data: &[u8],
data_pad: usize,
timeout: std::time::Duration,
bypass_readiness: bool,
) -> Result<Bytes> {
if !bypass_readiness {
self.wait_until_ready().await?;
}
let (tx, rx) = oneshot::channel();
self.pending
.lock()
.map_err(|_| NfsError::Rpc("pending map lock poisoned".to_string()))?
.insert(xid, tx);
let _pending_guard = PendingRequestGuard {
pending: Arc::clone(&self.pending),
xid,
};
let write_result = {
let mut writer = self.writer.lock().await;
async {
writer.write_all(header).await?;
if !data.is_empty() {
writer.write_all(data).await?;
if data_pad > 0 {
writer.write_all(&[0u8; 4][..data_pad]).await?;
}
}
Ok::<(), NfsError>(())
}
.await
};
write_result?;
match tokio::time::timeout(timeout, rx).await {
Ok(Ok(result)) => result,
Ok(Err(_)) => Err(NfsError::Io(std::io::Error::new(
std::io::ErrorKind::BrokenPipe,
"reader task terminated",
))),
Err(_) => Err(NfsError::Io(std::io::Error::new(
std::io::ErrorKind::TimedOut,
"RPC response timeout",
))),
}
}
async fn reconnect(self: &Arc<Self>, failed_gen: u64) -> Result<()> {
if self.shutdown_flag.load(Ordering::Acquire) {
return Err(NfsError::Io(std::io::Error::new(
std::io::ErrorKind::NotConnected,
"mux is shut down",
)));
}
let _reconnect = self.reconnect_lock.lock().await;
info!(addr = %self.addr, failed_gen, "initiating reconnection");
let current_gen = self.generation.load(Ordering::Acquire);
if current_gen > failed_gen {
debug!(addr = %self.addr, current_gen, failed_gen, "reconnection already performed by another caller");
return Ok(());
}
let stream = crate::connect_to_target(&self.addr, self.noresvport).await?;
let (reader, new_writer) = stream.into_split();
let reader = BufReader::with_capacity(1_048_576, reader);
let mut writer = self.writer.lock().await;
let current_gen = self.generation.load(Ordering::Acquire);
if current_gen > failed_gen {
debug!(addr = %self.addr, current_gen, failed_gen, "reconnection already performed by another caller (after connect)");
return Ok(()); }
if self.shutdown_flag.load(Ordering::Acquire) {
return Err(NfsError::Io(std::io::Error::new(
std::io::ErrorKind::NotConnected,
"mux is shut down",
)));
}
self.readiness
.store(CONNECTION_REBINDING, Ordering::Release);
let publication = RebindPublicationGuard {
readiness: &self.readiness,
notify: &self.readiness_notify,
published: false,
};
if let Ok(mut guard) = self.reader_handle.lock()
&& let Some(handle) = guard.take()
{
handle.abort();
}
{
let mut map = self
.pending
.lock()
.map_err(|_| NfsError::Rpc("pending map lock poisoned".to_string()))?;
if !map.is_empty() {
debug!(addr = %self.addr, pending_count = map.len(), "failing pending requests due to reconnection");
}
for (_, tx) in map.drain() {
let _ = tx.send(Err(NfsError::Io(std::io::Error::new(
std::io::ErrorKind::BrokenPipe,
"reconnecting",
))));
}
}
*writer = new_writer;
{
let mut guard = self
.reader_handle
.lock()
.map_err(|_| NfsError::Rpc("reader_handle lock poisoned".to_string()))?;
*guard = Some(tokio::spawn(reader_loop(
reader,
Arc::clone(&self.pending),
Arc::clone(&self.writer),
Arc::clone(&self.backchannel),
)));
}
drop(writer);
let next_generation = failed_gen.saturating_add(1);
let handler = self
.reconnect_handler
.lock()
.map_err(|_| NfsError::Rpc("reconnect handler lock poisoned".to_string()))?
.clone();
if let Some(handler) = handler
&& let Err(error) = handler(Client::new(Arc::clone(self), None), next_generation).await
{
return Err(error);
}
self.generation.store(next_generation, Ordering::Release);
publication.publish();
info!(addr = %self.addr, generation = next_generation, "reconnection ready");
Ok(())
}
async fn shutdown(&self) {
self.shutdown_flag.store(true, Ordering::Release);
debug!(addr = %self.addr, "shutting down StreamMux");
if let Ok(mut guard) = self.reader_handle.lock()
&& let Some(handle) = guard.take()
{
handle.abort();
}
let mut writer = self.writer.lock().await;
let _ = writer.shutdown().await;
if let Ok(mut map) = self.pending.lock() {
for (_, tx) in map.drain() {
let _ = tx.send(Err(NfsError::Io(std::io::Error::new(
std::io::ErrorKind::BrokenPipe,
"shutdown",
))));
}
}
}
}
impl Drop for StreamMux {
fn drop(&mut self) {
if let Ok(mut guard) = self.reader_handle.lock()
&& let Some(handle) = guard.take()
{
handle.abort();
}
if let Ok(mut map) = self.pending.lock() {
for (_, tx) in map.drain() {
let _ = tx.send(Err(NfsError::Io(std::io::Error::new(
std::io::ErrorKind::BrokenPipe,
"connection closed",
))));
}
}
}
}
async fn reader_loop(
mut reader: BufReader<OwnedReadHalf>,
pending: PendingMap,
writer: Arc<TokioMutex<OwnedWriteHalf>>,
backchannel: BackchannelSlot,
) {
loop {
match read_one_response(&mut reader).await {
Ok((xid, data)) => {
let msg_type = if data.len() >= 8 {
BigEndian::read_u32(&data[4..8])
} else {
MessageType::Response as u32
};
if msg_type == MessageType::Request as u32 {
dispatch_backchannel_call(xid, data, &writer, &backchannel).await;
continue;
}
match pending.lock() {
Ok(mut map) => match map.remove(&xid) {
Some(tx) => {
let _ = tx.send(Ok(data));
}
_ => {
debug!(
xid,
"dropping response for unmatched XID (likely stale retry)"
);
}
},
_ => {
warn!("pending map lock poisoned in reader loop, terminating");
break;
}
}
}
Err(e) => {
warn!(error = %e, "reader loop terminated due to connection error");
if let Ok(mut map) = pending.lock() {
for (_, tx) in map.drain() {
let _ = tx.send(Err(NfsError::Io(std::io::Error::new(
std::io::ErrorKind::BrokenPipe,
e.to_string(),
))));
}
}
break;
}
}
}
}
async fn dispatch_backchannel_call(
xid: u32,
data: Bytes,
writer: &Arc<TokioMutex<OwnedWriteHalf>>,
backchannel: &BackchannelSlot,
) {
let handler = match backchannel.lock() {
Ok(slot) => slot.clone(),
Err(_) => {
warn!("backchannel slot lock poisoned, dropping backchannel CALL");
return;
}
};
let Some(handler) = handler else {
debug!(
xid,
"backchannel CALL received but no handler registered, dropping"
);
return;
};
let Some(reply) = handler(data) else {
debug!(xid, "backchannel handler dropped CALL (parse error)");
return;
};
let mark = (reply.len() as u32) | 0x80000000;
let mut out = Vec::with_capacity(4 + reply.len());
out.extend_from_slice(&mark.to_be_bytes());
out.extend_from_slice(&reply);
let mut w = writer.lock().await;
if let Err(e) = w.write_all(&out).await {
warn!(xid, error = %e, "failed to write backchannel reply");
}
}
const MAX_RPC_RESPONSE: usize = 4 * 1024 * 1024 + 4096;
async fn read_one_response(reader: &mut BufReader<OwnedReadHalf>) -> Result<(u32, Bytes)> {
let mut hdr = [0u8; 4];
reader.read_exact(&mut hdr).await?;
let raw = BigEndian::read_u32(&hdr);
let last = (raw & 0x80000000) != 0;
let sz = (raw & 0x7FFFFFFF) as usize;
if sz > MAX_RPC_RESPONSE {
return Err(NfsError::Rpc(format!(
"RPC fragment size {} exceeds maximum {}",
sz, MAX_RPC_RESPONSE
)));
}
let mut buf = BytesMut::with_capacity(sz);
buf.resize(sz, 0);
reader.read_exact(&mut buf[..sz]).await?;
if !last {
loop {
reader.read_exact(&mut hdr).await?;
let raw = BigEndian::read_u32(&hdr);
let last = (raw & 0x80000000) != 0;
let sz = (raw & 0x7FFFFFFF) as usize;
let total = buf.len() + sz;
if total > MAX_RPC_RESPONSE {
return Err(NfsError::Rpc(format!(
"RPC accumulated response size {} exceeds maximum {}",
total, MAX_RPC_RESPONSE
)));
}
let offset = buf.len();
buf.resize(total, 0);
reader.read_exact(&mut buf[offset..]).await?;
if last {
break;
}
}
}
let xid = BigEndian::read_u32(&buf[0..4]);
Ok((xid, buf.freeze()))
}
#[derive(Debug, Clone)]
pub(crate) struct Client {
nfs_mux: Arc<StreamMux>,
mount_mux: Option<Arc<StreamMux>>,
}
impl std::fmt::Debug for StreamMux {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("StreamMux")
.field("addr", &self.addr)
.field("generation", &self.generation.load(Ordering::Acquire))
.finish()
}
}
impl Client {
pub(crate) fn new(nfs_mux: Arc<StreamMux>, mount_mux: Option<Arc<StreamMux>>) -> Self {
Self { nfs_mux, mount_mux }
}
pub(crate) fn enable_backchannel(&self, handler: BackchannelHandler) {
self.nfs_mux.enable_backchannel(handler);
}
pub(crate) fn set_reconnect_handler<F, Fut>(&self, handler: F) -> Result<()>
where
F: Fn(Client, u64) -> Fut + Send + Sync + 'static,
Fut: Future<Output = Result<()>> + Send + 'static,
{
self.nfs_mux
.set_reconnect_handler(Arc::new(move |client, generation| {
Box::pin(handler(client, generation))
}))
}
fn get_mux(&self, program: u32) -> Result<&Arc<StreamMux>> {
match program {
MOUNT_PROG => Ok(self.mount_mux.as_ref().unwrap_or(&self.nfs_mux)),
NFS_PROG | PORTMAP_PROG => Ok(&self.nfs_mux),
_ => Err(NfsError::InvalidInput(format!(
"unknown RPC program {}",
program
))),
}
}
pub(crate) async fn call(
&self,
msg_body: Vec<u8>,
replay_policy: ReplayPolicy,
timeout: std::time::Duration,
) -> Result<Bytes> {
self.call_with_data(msg_body, Bytes::new(), replay_policy, timeout)
.await
}
pub(crate) async fn call_with_data(
&self,
msg_body: Vec<u8>,
data: Bytes,
replay_policy: ReplayPolicy,
timeout: std::time::Duration,
) -> Result<Bytes> {
self.call_with_data_inner(msg_body, data, replay_policy, timeout, false)
.await
}
pub(crate) async fn call_during_reconnect(
&self,
msg_body: Vec<u8>,
timeout: std::time::Duration,
) -> Result<Bytes> {
self.call_with_data_inner(
msg_body,
Bytes::new(),
ReplayPolicy::ONE_ATTEMPT,
timeout,
true,
)
.await
}
async fn call_with_data_inner(
&self,
mut msg_body: Vec<u8>,
data: Bytes,
replay_policy: ReplayPolicy,
timeout: std::time::Duration,
bypass_readiness: bool,
) -> Result<Bytes> {
const SIZE_HDR_BIT: u32 = 0x80000000;
const PREFIX_LEN: usize = 12;
let max_attempts = replay_policy.max_attempts();
let mut attempt = 0usize;
let mut last_error = None;
let start = tokio::time::Instant::now();
let max_total = timeout.saturating_mul(3);
let program = if msg_body.len() >= 8 {
BigEndian::read_u32(&msg_body[4..8])
} else {
NFS_PROG
};
let mux = self.get_mux(program)?;
let data_len = data.len();
let data_pad = (4 - data_len % 4) % 4;
let payload_len = (8 + msg_body.len() + data_len + data_pad) as u32;
msg_body.splice(0..0, [0u8; PREFIX_LEN]);
BigEndian::write_u32(&mut msg_body[0..4], payload_len | SIZE_HDR_BIT);
BigEndian::write_u32(&mut msg_body[8..12], MessageType::Request as u32);
while attempt < max_attempts {
if start.elapsed() > max_total {
break;
}
let xid = get_xid();
BigEndian::write_u32(&mut msg_body[4..8], xid);
debug!(
xid,
attempt = attempt + 1,
max_attempts,
?replay_policy,
program,
"sending RPC request"
);
let r#gen = mux.generation();
let res = mux
.send_and_receive_inner(xid, &msg_body, &data, data_pad, timeout, bypass_readiness)
.await;
match res {
Ok(response_data) => {
trace!(xid, "RPC response received");
return parse_rpc_response(response_data, xid);
}
Err(e) => {
if bypass_readiness {
return Err(e);
}
let (is_conn_error, is_timeout) = match &e {
NfsError::Io(io_err) => (
matches!(
io_err.kind(),
std::io::ErrorKind::BrokenPipe
| std::io::ErrorKind::ConnectionAborted
| std::io::ErrorKind::ConnectionReset
),
io_err.kind() == std::io::ErrorKind::TimedOut,
),
_ => (false, false),
};
if is_conn_error {
attempt += 1;
if attempt >= max_attempts {
last_error = Some(e);
continue;
}
warn!(
xid,
attempt,
max_attempts,
error = %e,
"RPC call failed (connection error), reconnecting"
);
let jitter = rand::random_range(0..50u64);
let backoff = std::cmp::min(100u64 << (attempt - 1), 2000) + jitter;
tokio::time::sleep(tokio::time::Duration::from_millis(backoff)).await;
if let Err(reconn_err) = mux.reconnect(r#gen).await {
warn!(error = %reconn_err, "reconnect failed, will retry");
}
last_error = Some(e);
continue;
} else if is_timeout {
attempt += 1;
if attempt >= max_attempts {
last_error = Some(e);
continue;
}
warn!(
xid,
attempt,
max_attempts,
error = %e,
"RPC call timed out, retrying without reconnect"
);
last_error = Some(e);
continue;
} else {
error!(xid, error = %e, "RPC call failed with non-retryable error");
return Err(e);
}
}
}
}
error!(
max_attempts,
?replay_policy,
elapsed_ms = start.elapsed().as_millis() as u64,
program,
"RPC retries exhausted, giving up"
);
Err(last_error.unwrap_or_else(|| {
NfsError::Io(std::io::Error::new(
std::io::ErrorKind::TimedOut,
"RPC total attempt budget exhausted",
))
}))
}
pub(crate) async fn shutdown(&self) {
self.nfs_mux.shutdown().await;
if let Some(ref mount_mux) = self.mount_mux {
mount_mux.shutdown().await;
}
}
#[cfg(test)]
pub(crate) async fn new_dummy() -> Self {
let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
let addr = listener.local_addr().unwrap();
let (stream_result, _accept_result) =
tokio::join!(tokio::net::TcpStream::connect(addr), listener.accept());
let stream = stream_result.unwrap();
stream.set_nodelay(true).unwrap();
let (reader, writer) = stream.into_split();
let pending: PendingMap = Arc::new(std::sync::Mutex::new(HashMap::new()));
let writer = Arc::new(TokioMutex::new(writer));
let backchannel: BackchannelSlot = Arc::new(std::sync::Mutex::new(None));
let reader = BufReader::with_capacity(1_048_576, reader);
let reader_handle = tokio::spawn(reader_loop(
reader,
Arc::clone(&pending),
Arc::clone(&writer),
Arc::clone(&backchannel),
));
let mux = Arc::new(StreamMux {
writer,
pending,
backchannel,
addr,
noresvport: false,
generation: AtomicU64::new(0),
reconnect_lock: TokioMutex::new(()),
reconnect_handler: std::sync::Mutex::new(None),
readiness: AtomicU8::new(CONNECTION_READY),
readiness_notify: Notify::new(),
reader_handle: std::sync::Mutex::new(Some(reader_handle)),
shutdown_flag: AtomicBool::new(false),
});
Self {
nfs_mux: mux,
mount_mux: None,
}
}
}
fn parse_rpc_response(res: Bytes, xid: u32) -> Result<Bytes> {
let read_u32 = |data: &[u8], p: usize| -> Result<u32> {
if p + 4 > data.len() {
return Err(NfsError::Rpc("response truncated".to_string()));
}
Ok(BigEndian::read_u32(&data[p..p + 4]))
};
if res.len() < 8 {
error!(xid, response_len = res.len(), "RPC response too short");
return Err(NfsError::Rpc("response too short".to_string()));
}
let res_xid = BigEndian::read_u32(&res[0..4]);
let res_msgtype = BigEndian::read_u32(&res[4..8]);
if res_xid != xid {
error!(
expected_xid = xid,
actual_xid = res_xid,
"RPC response XID mismatch"
);
return Err(NfsError::Rpc(
"response id does not match expected one".to_string(),
));
}
if res_msgtype != MessageType::Response as u32 {
error!(
xid,
msgtype = res_msgtype,
"RPC response has unexpected message type"
);
return Err(NfsError::Rpc(
"response type does not match expected one".to_string(),
));
}
let mut pos = 8usize;
let msg_status = read_u32(&res, pos)? as i32;
pos += 4;
if msg_status != MessageStatus::Accepted as i32 {
error!(xid, msg_status, "RPC response rejected (bad status)");
return Err(NfsError::Rpc(
"could not parse response due to bad status".to_string(),
));
}
pos += 4; let verf_len = read_u32(&res, pos)? as usize;
pos += 4;
let verf_padded = verf_len + (4 - verf_len % 4) % 4;
if pos + verf_padded > res.len() {
error!(
xid,
response_len = res.len(),
"RPC response truncated (verifier)"
);
return Err(NfsError::Rpc("response truncated (verifier)".to_string()));
}
pos += verf_padded;
let accept_status = read_u32(&res, pos)? as i32;
pos += 4;
if accept_status != AcceptStatus::Success as i32 {
error!(xid, accept_status, "RPC request rejected by server");
return Err(NfsError::Rpc("request rejected".to_string()));
}
Ok(res.slice(pos..))
}
#[derive(Debug, Clone, PartialEq)]
enum MessageType {
Request = 0,
Response = 1,
}
enum MessageStatus {
Accepted = 0,
#[allow(unused)]
Denied = 1,
}
enum AcceptStatus {
Success = 0,
#[allow(unused)]
ProgUnavail = 1,
#[allow(unused)]
ProgMismatch = 2,
#[allow(unused)]
ProcUnavail = 3,
#[allow(unused)]
GarbageArgs = 4,
}
static XID: AtomicU32 = AtomicU32::new(0);
fn get_xid() -> u32 {
if XID.load(Ordering::Relaxed) == 0 {
XID.compare_exchange(0, get_current_time(), Ordering::Relaxed, Ordering::Relaxed)
.ok();
}
XID.fetch_add(1, Ordering::Relaxed).wrapping_add(1)
}
pub(crate) fn get_current_time() -> u32 {
let now = std::time::SystemTime::now();
let since_epoch = now
.duration_since(std::time::UNIX_EPOCH)
.unwrap_or_default();
(since_epoch.as_secs() as u32).wrapping_mul(1000) + since_epoch.subsec_millis()
}
#[cfg(test)]
mod tests {
use super::*;
use std::sync::atomic::AtomicUsize;
async fn read_test_record(stream: &mut tokio::net::TcpStream) -> std::io::Result<Vec<u8>> {
let marker = stream.read_u32().await?;
let len = (marker & 0x7fff_ffff) as usize;
let mut record = vec![0; len];
stream.read_exact(&mut record).await?;
Ok(record)
}
async fn write_test_rpc_reply(
stream: &mut tokio::net::TcpStream,
xid: u32,
payload: &[u8],
) -> std::io::Result<()> {
let mut reply = Vec::with_capacity(24 + payload.len());
reply.extend_from_slice(&xid.to_be_bytes());
reply.extend_from_slice(&(MessageType::Response as u32).to_be_bytes());
reply.extend_from_slice(&(MessageStatus::Accepted as u32).to_be_bytes());
reply.extend_from_slice(&0u32.to_be_bytes()); reply.extend_from_slice(&0u32.to_be_bytes()); reply.extend_from_slice(&(AcceptStatus::Success as u32).to_be_bytes());
stream
.write_u32(0x8000_0000 | reply.len().saturating_add(payload.len()) as u32)
.await?;
stream.write_all(&reply).await?;
stream.write_all(payload).await
}
#[tokio::test]
async fn concurrent_reconnect_observers_run_one_effective_rebind() {
let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
let addr = listener.local_addr().unwrap();
let accepted = Arc::new(Notify::new());
let accepted_for_server = Arc::clone(&accepted);
let server = tokio::spawn(async move {
let (first, _) = listener.accept().await?;
let (second, _) = listener.accept().await?;
accepted_for_server.notify_one();
tokio::time::sleep(std::time::Duration::from_secs(5)).await;
Ok::<_, std::io::Error>((first, second))
});
let mux = StreamMux::connect(addr, true).await.unwrap();
let binds = Arc::new(AtomicUsize::new(0));
let binds_for_handler = Arc::clone(&binds);
mux.set_reconnect_handler(Arc::new(move |_client, generation| {
let binds = Arc::clone(&binds_for_handler);
Box::pin(async move {
assert_eq!(generation, 1);
binds.fetch_add(1, Ordering::AcqRel);
Ok(())
})
}))
.unwrap();
let tasks = (0..64)
.map(|_| {
let mux = Arc::clone(&mux);
tokio::spawn(async move { mux.reconnect(0).await })
})
.collect::<Vec<_>>();
accepted.notified().await;
for task in tasks {
task.await.unwrap().unwrap();
}
assert_eq!(binds.load(Ordering::Acquire), 1);
assert_eq!(mux.generation(), 1);
assert_eq!(mux.readiness.load(Ordering::Acquire), CONNECTION_READY);
server.abort();
}
#[tokio::test]
async fn failed_rebind_marks_replacement_connection_unready() {
let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
let addr = listener.local_addr().unwrap();
let server = tokio::spawn(async move {
let (first, _) = listener.accept().await?;
let (second, _) = listener.accept().await?;
tokio::time::sleep(std::time::Duration::from_secs(5)).await;
Ok::<_, std::io::Error>((first, second))
});
let mux = StreamMux::connect(addr, true).await.unwrap();
mux.set_reconnect_handler(Arc::new(|_client, _generation| {
Box::pin(async { Err(NfsError::Rpc("injected bind failure".to_string())) })
}))
.unwrap();
assert!(mux.reconnect(0).await.is_err());
assert_eq!(mux.readiness.load(Ordering::Acquire), CONNECTION_FAILED);
assert!(mux.wait_until_ready().await.is_err());
server.abort();
}
#[tokio::test]
async fn cancelled_rebind_marks_replacement_connection_unready() {
let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
let addr = listener.local_addr().unwrap();
let server = tokio::spawn(async move {
let (first, _) = listener.accept().await?;
let (second, _) = listener.accept().await?;
tokio::time::sleep(std::time::Duration::from_secs(5)).await;
Ok::<_, std::io::Error>((first, second))
});
let mux = StreamMux::connect(addr, true).await.unwrap();
let entered = Arc::new(Notify::new());
let entered_for_handler = Arc::clone(&entered);
mux.set_reconnect_handler(Arc::new(move |_client, _generation| {
let entered = Arc::clone(&entered_for_handler);
Box::pin(async move {
entered.notify_one();
std::future::pending::<Result<()>>().await
})
}))
.unwrap();
let mux_for_task = Arc::clone(&mux);
let task = tokio::spawn(async move { mux_for_task.reconnect(0).await });
entered.notified().await;
task.abort();
let _ = task.await;
assert_eq!(mux.readiness.load(Ordering::Acquire), CONNECTION_FAILED);
assert!(mux.wait_until_ready().await.is_err());
server.abort();
}
#[test]
fn message_rpc_version() {
let body = vec![0u8, 0, 0, 2, 0, 0, 0, 3, 0, 0, 0, 4, 0, 0, 0, 5];
assert_eq!(BigEndian::read_u32(&body[0..4]), 2);
}
#[test]
fn message_program() {
let body = vec![0u8, 0, 0, 2, 0, 0, 0, 3, 0, 0, 0, 4, 0, 0, 0, 5];
assert_eq!(BigEndian::read_u32(&body[4..8]), 3);
}
#[test]
fn message_version() {
let body = vec![0u8, 0, 0, 2, 0, 0, 0, 3, 0, 0, 0, 4, 0, 0, 0, 5];
assert_eq!(BigEndian::read_u32(&body[8..12]), 4);
}
#[test]
fn message_procedure() {
let body = vec![0u8, 0, 0, 2, 0, 0, 0, 3, 0, 0, 0, 4, 0, 0, 0, 5];
assert_eq!(BigEndian::read_u32(&body[12..16]), 5);
}
#[tokio::test]
async fn portmap_error_includes_underlying_detail() {
let dead_addr: SocketAddr = "127.0.0.1:1".parse().unwrap();
let auth = Auth::new_null();
let res = portmap(&vec![dead_addr], NFS_PROG, NFS3_VERSION, &auth, 2, false).await;
let err = res.expect_err("dead port should fail");
let msg = err.to_string();
assert!(
msg.contains("127.0.0.1:1")
|| msg.to_lowercase().contains("refused")
|| msg.to_lowercase().contains("connect"),
"portmap error should expose underlying detail, got: {}",
msg
);
}
#[tokio::test]
async fn retransmission_changes_only_xid_and_preserves_zero_copy_payload() {
let listener = tokio::net::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?;
let first = read_test_record(&mut stream).await?;
let second = read_test_record(&mut stream).await?;
if first.len() < 8 || second.len() < 8 {
return Err(std::io::Error::other("RPC request too short"));
}
if first[0..4] == second[0..4] {
return Err(std::io::Error::other(
"transport attempts must use distinct XIDs",
));
}
if first[4..] != second[4..] {
return Err(std::io::Error::other(
"logical request changed across retransmission",
));
}
let xid = BigEndian::read_u32(&second[0..4]);
write_test_rpc_reply(&mut stream, xid, b"cached-result").await?;
Ok::<(Vec<u8>, Vec<u8>), std::io::Error>((first, second))
});
let mux = StreamMux::connect(addr, true).await.unwrap();
let client = Client::new(mux, None);
let payload = Bytes::from(vec![0x5a; 64 * 1024]);
let payload_ptr = payload.as_ptr();
let session_id = [0x33; 16];
let state_id = [0x44; 16];
let builder = crate::nfs41::compound::CompoundBuilder::new("retry-write")
.sequence(&session_id, 17, 3, 7)
.putfh(b"file-handle")
.write_header(&state_id, 4096, 2, payload.len() as u32)
.apply_sequence_cache_policy(4096)
.unwrap();
let mut body = Vec::new();
builder.encode_with_header(&Auth::new_null(), &mut body);
let response = client
.call_with_data(
body,
payload.clone(),
ReplayPolicy::byte_identical(2),
std::time::Duration::from_millis(20),
)
.await
.unwrap();
assert_eq!(response, b"cached-result"[..]);
assert_eq!(payload.as_ptr(), payload_ptr);
let (first, second) = server.await.unwrap().unwrap();
let mut request_identity = Vec::from(session_id);
request_identity.extend_from_slice(&17u32.to_be_bytes());
request_identity.extend_from_slice(&3u32.to_be_bytes());
assert!(
first
.windows(request_identity.len())
.any(|window| window == request_identity)
);
assert_eq!(&first[first.len() - payload.len()..], payload.as_ref());
assert_eq!(&second[second.len() - payload.len()..], payload.as_ref());
client.shutdown().await;
}
#[tokio::test]
async fn one_attempt_does_not_retransmit_and_preserves_timeout() {
let listener = tokio::net::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?;
let first = read_test_record(&mut stream).await?;
let second = tokio::time::timeout(
std::time::Duration::from_millis(80),
read_test_record(&mut stream),
)
.await;
Ok::<(Vec<u8>, bool), std::io::Error>((first, second.is_ok()))
});
let mux = StreamMux::connect(addr, true).await.unwrap();
let client = Client::new(mux, None);
let mut body = Vec::new();
body.extend_from_slice(&RPC_VERSION.to_be_bytes());
body.extend_from_slice(&NFS_PROG.to_be_bytes());
body.extend_from_slice(&crate::nfs41::NFS4_VERSION.to_be_bytes());
body.extend_from_slice(&1u32.to_be_bytes());
let err = client
.call(
body,
ReplayPolicy::ONE_ATTEMPT,
std::time::Duration::from_millis(20),
)
.await
.expect_err("an unanswered one-attempt call must time out");
assert!(
matches!(err, NfsError::Io(ref error) if error.kind() == std::io::ErrorKind::TimedOut),
"one-attempt call must preserve its authoritative timeout: {err}"
);
let (first, saw_second) = server.await.unwrap().unwrap();
assert!(!first.is_empty());
assert!(!saw_second, "one-attempt policy must not retransmit");
client.shutdown().await;
}
#[tokio::test]
async fn one_attempt_connection_failure_does_not_reconnect() {
let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
let addr = listener.local_addr().unwrap();
let server = tokio::spawn(async move {
let (mut first_stream, _) = listener.accept().await?;
let first = read_test_record(&mut first_stream).await?;
drop(first_stream);
let reconnected =
tokio::time::timeout(std::time::Duration::from_millis(250), listener.accept())
.await
.is_ok();
Ok::<(Vec<u8>, bool), std::io::Error>((first, reconnected))
});
let mux = StreamMux::connect(addr, true).await.unwrap();
let client = Client::new(mux, None);
let mut body = Vec::new();
body.extend_from_slice(&RPC_VERSION.to_be_bytes());
body.extend_from_slice(&NFS_PROG.to_be_bytes());
body.extend_from_slice(&crate::nfs41::NFS4_VERSION.to_be_bytes());
body.extend_from_slice(&1u32.to_be_bytes());
let err = client
.call(
body,
ReplayPolicy::ONE_ATTEMPT,
std::time::Duration::from_secs(1),
)
.await
.expect_err("closed connection must fail the call");
assert!(matches!(err, NfsError::Io(_)));
let (first, reconnected) = server.await.unwrap().unwrap();
assert!(!first.is_empty());
assert!(!reconnected, "one-attempt policy must not reconnect");
client.shutdown().await;
}
#[tokio::test]
async fn byte_identical_replay_reconnects_once_after_connection_failure() {
let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
let addr = listener.local_addr().unwrap();
let server = tokio::spawn(async move {
let (mut first_stream, _) = listener.accept().await?;
let first = read_test_record(&mut first_stream).await?;
drop(first_stream);
let (mut second_stream, _) = listener.accept().await?;
let second = read_test_record(&mut second_stream).await?;
let xid = BigEndian::read_u32(&second[0..4]);
write_test_rpc_reply(&mut second_stream, xid, b"after-reconnect").await?;
Ok::<(Vec<u8>, Vec<u8>), std::io::Error>((first, second))
});
let mux = StreamMux::connect(addr, true).await.unwrap();
let client = Client::new(mux, None);
let mut body = Vec::new();
body.extend_from_slice(&RPC_VERSION.to_be_bytes());
body.extend_from_slice(&NFS_PROG.to_be_bytes());
body.extend_from_slice(&crate::nfs41::NFS4_VERSION.to_be_bytes());
body.extend_from_slice(&1u32.to_be_bytes());
let response = client
.call(
body,
ReplayPolicy::byte_identical(2),
std::time::Duration::from_secs(1),
)
.await
.unwrap();
assert_eq!(response, b"after-reconnect"[..]);
let (first, second) = server.await.unwrap().unwrap();
assert_ne!(&first[0..4], &second[0..4], "attempts need fresh XIDs");
assert_eq!(&first[4..], &second[4..], "logical request must be stable");
client.shutdown().await;
}
#[test]
#[should_panic(expected = "requires at least 2 attempts")]
fn byte_identical_replay_requires_multiple_attempts() {
let _ = ReplayPolicy::byte_identical(1);
}
#[tokio::test]
async fn cancelled_rpc_removes_pending_xid() {
let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
let addr = listener.local_addr().unwrap();
let (received_tx, received_rx) = oneshot::channel();
let server = tokio::spawn(async move {
let (mut stream, _) = listener.accept().await?;
let _request = read_test_record(&mut stream).await?;
let _ = received_tx.send(());
std::future::pending::<()>().await;
#[allow(unreachable_code)]
Ok::<(), std::io::Error>(())
});
let mux = StreamMux::connect(addr, true).await.unwrap();
let client = Client::new(Arc::clone(&mux), None);
let mut body = Vec::new();
body.extend_from_slice(&RPC_VERSION.to_be_bytes());
body.extend_from_slice(&NFS_PROG.to_be_bytes());
body.extend_from_slice(&crate::nfs41::NFS4_VERSION.to_be_bytes());
body.extend_from_slice(&1u32.to_be_bytes());
let call_client = client.clone();
let call = tokio::spawn(async move {
call_client
.call(
body,
ReplayPolicy::ONE_ATTEMPT,
std::time::Duration::from_secs(30),
)
.await
});
received_rx.await.unwrap();
call.abort();
let _ = call.await;
tokio::task::yield_now().await;
assert_eq!(mux.pending.lock().unwrap().len(), 0);
server.abort();
let _ = server.await;
client.shutdown().await;
}
}