use std::collections::HashMap;
use std::sync::{Arc, Mutex};
use bytes::{Buf, Bytes};
use tokio::sync::mpsc;
use tracing::{debug, warn};
use crate::error::{NfsError, Result};
use crate::rpc::BackchannelHandler;
const MAX_CB_OPS: u32 = 64;
const MAX_REFERRING_LISTS: usize = 256;
pub(crate) const CB_PROGRAM: u32 = 0x40000000;
const NFS4ERR_SEQ_MISORDERED: u32 = 10063;
#[derive(Debug, Clone)]
pub(crate) enum RecallNotification {
Delegation {
stateid: [u8; 16],
#[allow(dead_code)]
truncate: bool,
fh: Bytes,
},
LayoutFile {
stateid: [u8; 16],
fh: Bytes,
offset: u64,
length: u64,
iomode: u32,
},
LayoutAll,
}
pub(crate) fn make_backchannel_handler(
session_id: [u8; 16],
recall_tx: mpsc::Sender<RecallNotification>,
) -> BackchannelHandler {
let slot_seqs: Arc<Mutex<HashMap<u32, u32>>> = Arc::new(Mutex::new(HashMap::new()));
Arc::new(move |frame: Bytes| {
let mut buf = frame;
handle_cb_compound(&mut buf, &session_id, &recall_tx, &slot_seqs)
})
}
fn handle_cb_compound(
buf: &mut Bytes,
session_id: &[u8; 16],
recall_tx: &mpsc::Sender<RecallNotification>,
slot_seqs: &Mutex<HashMap<u32, u32>>,
) -> Option<Vec<u8>> {
match parse_cb_compound(buf, session_id, recall_tx, slot_seqs) {
Ok(reply) => Some(reply),
Err(e) => {
warn!(error = %e, "failed to handle CB_COMPOUND");
None
}
}
}
fn parse_cb_compound(
buf: &mut Bytes,
session_id: &[u8; 16],
recall_tx: &mpsc::Sender<RecallNotification>,
slot_seqs: &Mutex<HashMap<u32, u32>>,
) -> Result<Vec<u8>> {
if buf.remaining() < 24 {
return Err(NfsError::Xdr("CB RPC header too short".to_string()));
}
let xid = buf.get_u32();
let _msg_type = buf.get_u32(); let _rpc_vers = buf.get_u32();
let _program = buf.get_u32();
let _version = buf.get_u32();
let procedure = buf.get_u32();
skip_rpc_auth(buf)?;
skip_rpc_auth(buf)?;
if procedure != 1 {
return Ok(build_rpc_reply(xid, &[]));
}
let _tag = skip_opaque(buf)?;
if buf.remaining() < 8 {
return Err(NfsError::Xdr("CB_COMPOUND args truncated".to_string()));
}
let _minor_version = buf.get_u32();
let _callback_ident = buf.get_u32();
if buf.remaining() < 4 {
return Err(NfsError::Xdr("CB_COMPOUND ops count truncated".to_string()));
}
let num_ops = buf.get_u32();
if num_ops > MAX_CB_OPS {
return Err(NfsError::Xdr(format!(
"CB_COMPOUND has {} ops, max {}",
num_ops, MAX_CB_OPS
)));
}
let mut reply_ops = Vec::new();
for _ in 0..num_ops {
if buf.remaining() < 4 {
break;
}
let opcode = buf.get_u32();
match opcode {
11 => {
if buf.remaining() < 32 {
break;
}
let mut cb_session_id = [0u8; 16];
buf.copy_to_slice(&mut cb_session_id);
let cb_sequenceid = buf.get_u32();
let cb_slotid = buf.get_u32();
let cb_highest_slotid = buf.get_u32();
let _cachethis = buf.get_u32();
if buf.remaining() >= 4 {
let n = buf.get_u32() as usize;
if n > MAX_REFERRING_LISTS {
return Err(NfsError::Xdr(format!(
"too many referring_call_lists: {}",
n
)));
}
for _ in 0..n {
if buf.remaining() < 16 {
return Err(NfsError::Xdr(
"referring_call sessionid truncated".to_string(),
));
}
buf.advance(16);
if buf.remaining() < 4 {
return Err(NfsError::Xdr(
"referring_call count truncated".to_string(),
));
}
let m = buf.get_u32() as usize;
let needed = m
.checked_mul(8)
.ok_or_else(|| NfsError::Xdr("referring_call overflow".to_string()))?;
if buf.remaining() < needed {
return Err(NfsError::Xdr("referring_call data truncated".to_string()));
}
buf.advance(needed);
}
}
if cb_session_id != *session_id {
warn!("CB_SEQUENCE session ID mismatch, ignoring");
return Err(NfsError::Xdr("CB_SEQUENCE session ID mismatch".to_string()));
}
let expected_seq = {
let map = slot_seqs
.lock()
.map_err(|_| NfsError::Rpc("cb slot table lock poisoned".to_string()))?;
map.get(&cb_slotid).copied().unwrap_or(1)
};
if cb_sequenceid != expected_seq {
warn!(
cb_slotid,
cb_sequenceid, expected_seq, "CB_SEQUENCE misordered — rejecting"
);
let mut op_reply = Vec::new();
op_reply.extend_from_slice(&opcode.to_be_bytes());
op_reply.extend_from_slice(&NFS4ERR_SEQ_MISORDERED.to_be_bytes());
reply_ops.push(op_reply);
break; }
{
let mut map = slot_seqs
.lock()
.map_err(|_| NfsError::Rpc("cb slot table lock poisoned".to_string()))?;
map.insert(cb_slotid, cb_sequenceid.wrapping_add(1));
}
let mut op_reply = Vec::new();
op_reply.extend_from_slice(&opcode.to_be_bytes());
op_reply.extend_from_slice(&0u32.to_be_bytes()); op_reply.extend_from_slice(&cb_session_id);
op_reply.extend_from_slice(&cb_sequenceid.to_be_bytes());
op_reply.extend_from_slice(&cb_slotid.to_be_bytes());
op_reply.extend_from_slice(&cb_highest_slotid.to_be_bytes());
op_reply.extend_from_slice(&cb_highest_slotid.to_be_bytes()); reply_ops.push(op_reply);
debug!(
slotid = cb_slotid,
seqid = cb_sequenceid,
"CB_SEQUENCE handled"
);
}
4 => {
if buf.remaining() < 20 {
break;
}
let mut stateid = [0u8; 16];
buf.copy_to_slice(&mut stateid);
let truncate = buf.get_u32() != 0;
let fh = read_opaque(buf)?;
debug!(fh_len = fh.len(), truncate, "CB_RECALL received");
if let Err(e) = recall_tx.try_send(RecallNotification::Delegation {
stateid,
truncate,
fh,
}) {
warn!("CB_RECALL notification channel full or closed: {}", e);
}
let mut op_reply = Vec::new();
op_reply.extend_from_slice(&opcode.to_be_bytes());
op_reply.extend_from_slice(&0u32.to_be_bytes()); reply_ops.push(op_reply);
}
5 => {
if buf.remaining() < 16 {
break;
}
let _layout_type = buf.get_u32();
let iomode = buf.get_u32();
let _changed = buf.get_u32();
let recalltype = buf.get_u32();
let notification = match recalltype {
1 => {
let fh = read_opaque(buf)?;
if buf.remaining() < 32 {
return Err(NfsError::Xdr(
"CB_LAYOUTRECALL file args truncated".to_string(),
));
}
let offset = buf.get_u64();
let length = buf.get_u64();
let mut stateid = [0u8; 16];
buf.copy_to_slice(&mut stateid);
Some(RecallNotification::LayoutFile {
stateid,
fh,
offset,
length,
iomode,
})
}
2 => {
if buf.remaining() < 16 {
return Err(NfsError::Xdr(
"CB_LAYOUTRECALL fsid truncated".to_string(),
));
}
buf.advance(16);
Some(RecallNotification::LayoutAll)
}
3 => Some(RecallNotification::LayoutAll),
_ => None,
};
match notification {
Some(n) => {
debug!(recalltype, iomode, "CB_LAYOUTRECALL received");
if let Err(e) = recall_tx.try_send(n) {
warn!("CB_LAYOUTRECALL notification channel full or closed: {}", e);
}
let mut op_reply = Vec::new();
op_reply.extend_from_slice(&opcode.to_be_bytes());
op_reply.extend_from_slice(&0u32.to_be_bytes()); reply_ops.push(op_reply);
}
None => {
let mut op_reply = Vec::new();
op_reply.extend_from_slice(&opcode.to_be_bytes());
op_reply.extend_from_slice(&10022u32.to_be_bytes()); reply_ops.push(op_reply);
debug!(recalltype, "CB_LAYOUTRECALL unknown recalltype");
break; }
}
}
_ => {
let mut op_reply = Vec::new();
op_reply.extend_from_slice(&opcode.to_be_bytes());
op_reply.extend_from_slice(&10044u32.to_be_bytes()); reply_ops.push(op_reply);
debug!(opcode, "unknown callback op, returning OP_ILLEGAL");
}
}
}
let mut compound_res = Vec::new();
compound_res.extend_from_slice(&0u32.to_be_bytes()); compound_res.extend_from_slice(&0u32.to_be_bytes()); compound_res.extend_from_slice(&(reply_ops.len() as u32).to_be_bytes());
for op in &reply_ops {
compound_res.extend_from_slice(op);
}
Ok(build_rpc_reply(xid, &compound_res))
}
fn build_rpc_reply(xid: u32, body: &[u8]) -> Vec<u8> {
let mut reply = Vec::with_capacity(24 + body.len());
reply.extend_from_slice(&xid.to_be_bytes()); reply.extend_from_slice(&1u32.to_be_bytes()); reply.extend_from_slice(&0u32.to_be_bytes()); reply.extend_from_slice(&0u32.to_be_bytes()); reply.extend_from_slice(&0u32.to_be_bytes()); reply.extend_from_slice(&0u32.to_be_bytes()); reply.extend_from_slice(body);
reply
}
fn skip_rpc_auth(buf: &mut Bytes) -> Result<()> {
if buf.remaining() < 8 {
return Err(NfsError::Xdr("RPC auth truncated".to_string()));
}
let _flavor = buf.get_u32();
let len = buf.get_u32() as usize;
let padded = (len + 3) & !3;
if buf.remaining() < padded {
return Err(NfsError::Xdr("RPC auth body truncated".to_string()));
}
buf.advance(padded);
Ok(())
}
fn skip_opaque(buf: &mut Bytes) -> Result<usize> {
if buf.remaining() < 4 {
return Err(NfsError::Xdr("opaque length truncated".to_string()));
}
let len = buf.get_u32() as usize;
let padded = (len + 3) & !3;
if buf.remaining() < padded {
return Err(NfsError::Xdr("opaque data truncated".to_string()));
}
buf.advance(padded);
Ok(len)
}
fn read_opaque(buf: &mut Bytes) -> Result<Bytes> {
if buf.remaining() < 4 {
return Err(NfsError::Xdr("opaque length truncated".to_string()));
}
let len = buf.get_u32() as usize;
let padded = (len + 3) & !3;
if buf.remaining() < padded {
return Err(NfsError::Xdr("opaque data truncated".to_string()));
}
let data = buf.slice(..len);
buf.advance(padded);
Ok(data)
}
#[cfg(test)]
#[allow(clippy::unwrap_used)]
mod tests {
use super::*;
fn be(v: u32) -> [u8; 4] {
v.to_be_bytes()
}
fn u32_at(b: &[u8], off: usize) -> u32 {
u32::from_be_bytes([b[off], b[off + 1], b[off + 2], b[off + 3]])
}
fn cb_sequence_frame(xid: u32, session_id: &[u8; 16], slotid: u32, seqid: u32) -> Bytes {
let mut f = Vec::new();
f.extend_from_slice(&be(xid)); f.extend_from_slice(&be(0)); f.extend_from_slice(&be(2)); f.extend_from_slice(&be(CB_PROGRAM)); f.extend_from_slice(&be(1)); f.extend_from_slice(&be(1)); f.extend_from_slice(&be(0)); f.extend_from_slice(&be(0)); f.extend_from_slice(&be(0)); f.extend_from_slice(&be(0)); f.extend_from_slice(&be(0)); f.extend_from_slice(&be(1)); f.extend_from_slice(&be(0)); f.extend_from_slice(&be(1)); f.extend_from_slice(&be(11)); f.extend_from_slice(session_id); f.extend_from_slice(&be(seqid)); f.extend_from_slice(&be(slotid)); f.extend_from_slice(&be(slotid)); f.extend_from_slice(&be(0)); f.extend_from_slice(&be(0)); Bytes::from(f)
}
fn cb_recall_frame(xid: u32, stateid: &[u8; 16], truncate: bool, fh: &[u8]) -> Bytes {
let mut f = Vec::new();
f.extend_from_slice(&be(xid));
f.extend_from_slice(&be(0)); f.extend_from_slice(&be(2));
f.extend_from_slice(&be(CB_PROGRAM));
f.extend_from_slice(&be(1));
f.extend_from_slice(&be(1)); f.extend_from_slice(&be(0)); f.extend_from_slice(&be(0));
f.extend_from_slice(&be(0)); f.extend_from_slice(&be(0));
f.extend_from_slice(&be(0)); f.extend_from_slice(&be(1)); f.extend_from_slice(&be(0)); f.extend_from_slice(&be(1)); f.extend_from_slice(&be(4)); f.extend_from_slice(stateid); f.extend_from_slice(&be(if truncate { 1 } else { 0 })); f.extend_from_slice(&be(fh.len() as u32)); f.extend_from_slice(fh);
let pad = (4 - fh.len() % 4) % 4;
f.extend_from_slice(&[0u8; 4][..pad]);
Bytes::from(f)
}
#[test]
fn cb_sequence_ok_then_misordered() {
let session_id = [7u8; 16];
let (tx, _rx) = mpsc::channel(4);
let slots = Mutex::new(HashMap::new());
let mut frame = cb_sequence_frame(0xAABBCCDD, &session_id, 0, 1);
let reply = handle_cb_compound(&mut frame, &session_id, &tx, &slots).unwrap();
assert_eq!(u32_at(&reply, 0), 0xAABBCCDD); assert_eq!(u32_at(&reply, 4), 1); assert_eq!(u32_at(&reply, 24), 0); assert_eq!(u32_at(&reply, 32), 1); assert_eq!(u32_at(&reply, 36), 11); assert_eq!(u32_at(&reply, 40), 0);
let mut frame2 = cb_sequence_frame(0xAABBCCDE, &session_id, 0, 1);
let reply2 = handle_cb_compound(&mut frame2, &session_id, &tx, &slots).unwrap();
assert_eq!(u32_at(&reply2, 36), 11); assert_eq!(u32_at(&reply2, 40), NFS4ERR_SEQ_MISORDERED); }
#[test]
fn cb_sequence_session_mismatch_dropped() {
let session_id = [1u8; 16];
let wrong = [2u8; 16];
let (tx, _rx) = mpsc::channel(4);
let slots = Mutex::new(HashMap::new());
let mut frame = cb_sequence_frame(1, &wrong, 0, 1);
assert!(handle_cb_compound(&mut frame, &session_id, &tx, &slots).is_none());
}
#[test]
fn cb_recall_forwards_notification() {
let session_id = [9u8; 16];
let stateid = [0xEE; 16];
let fh = b"file-handle-xyz"; let (tx, mut rx) = mpsc::channel(4);
let slots = Mutex::new(HashMap::new());
let mut frame = cb_recall_frame(0x11223344, &stateid, true, fh);
let reply = handle_cb_compound(&mut frame, &session_id, &tx, &slots).unwrap();
assert_eq!(u32_at(&reply, 36), 4); assert_eq!(u32_at(&reply, 40), 0);
match rx.try_recv().expect("recall notification forwarded") {
RecallNotification::Delegation {
stateid: sid,
truncate,
fh: nfh,
} => {
assert_eq!(sid, stateid);
assert!(truncate);
assert_eq!(&nfh[..], &fh[..]);
}
other => panic!("expected Delegation, got {other:?}"),
}
}
fn cb_layoutrecall_frame(
xid: u32,
iomode: u32,
recalltype: u32,
file_body: Option<(&[u8], u64, u64, &[u8; 16])>,
) -> Bytes {
let mut f = Vec::new();
f.extend_from_slice(&be(xid));
f.extend_from_slice(&be(0)); f.extend_from_slice(&be(2));
f.extend_from_slice(&be(CB_PROGRAM));
f.extend_from_slice(&be(1));
f.extend_from_slice(&be(1)); f.extend_from_slice(&be(0)); f.extend_from_slice(&be(0));
f.extend_from_slice(&be(0)); f.extend_from_slice(&be(0));
f.extend_from_slice(&be(0)); f.extend_from_slice(&be(1)); f.extend_from_slice(&be(0)); f.extend_from_slice(&be(1)); f.extend_from_slice(&be(5)); f.extend_from_slice(&be(1)); f.extend_from_slice(&be(iomode)); f.extend_from_slice(&be(0)); f.extend_from_slice(&be(recalltype));
match recalltype {
1 => {
let (fh, offset, length, stateid) = file_body.expect("FILE recall needs body");
f.extend_from_slice(&be(fh.len() as u32));
f.extend_from_slice(fh);
let pad = (4 - fh.len() % 4) % 4;
f.extend_from_slice(&[0u8; 4][..pad]);
f.extend_from_slice(&offset.to_be_bytes());
f.extend_from_slice(&length.to_be_bytes());
f.extend_from_slice(stateid);
}
2 => {
f.extend_from_slice(&[0u8; 16]); }
_ => {}
}
Bytes::from(f)
}
#[test]
fn cb_layoutrecall_file_forwards_notification() {
let session_id = [3u8; 16];
let stateid = [0xAB; 16];
let fh = b"layout-fh-123"; let (tx, mut rx) = mpsc::channel(4);
let slots = Mutex::new(HashMap::new());
let mut frame = cb_layoutrecall_frame(
0x55667788,
2, 1, Some((fh, 0, u64::MAX, &stateid)),
);
let reply = handle_cb_compound(&mut frame, &session_id, &tx, &slots).unwrap();
assert_eq!(u32_at(&reply, 36), 5); assert_eq!(u32_at(&reply, 40), 0);
match rx.try_recv().expect("layout recall forwarded") {
RecallNotification::LayoutFile {
stateid: sid,
fh: nfh,
offset,
length,
iomode,
} => {
assert_eq!(sid, stateid);
assert_eq!(&nfh[..], &fh[..]);
assert_eq!(offset, 0);
assert_eq!(length, u64::MAX);
assert_eq!(iomode, 2);
}
other => panic!("expected LayoutFile, got {other:?}"),
}
}
#[test]
fn cb_layoutrecall_all_forwards_notification() {
let session_id = [3u8; 16];
let (tx, mut rx) = mpsc::channel(4);
let slots = Mutex::new(HashMap::new());
for recalltype in [2u32, 3u32] {
let mut frame = cb_layoutrecall_frame(1, 3, recalltype, None);
let reply = handle_cb_compound(&mut frame, &session_id, &tx, &slots).unwrap();
assert_eq!(u32_at(&reply, 40), 0); match rx.try_recv().expect("layout recall forwarded") {
RecallNotification::LayoutAll => {}
other => panic!("expected LayoutAll, got {other:?}"),
}
}
}
#[test]
fn cb_layoutrecall_truncated_no_panic() {
let session_id = [3u8; 16];
let (tx, mut rx) = mpsc::channel(4);
let slots = Mutex::new(HashMap::new());
let full = cb_layoutrecall_frame(1, 2, 1, Some((b"fh", 7, 9, &[1u8; 16])));
let mut truncated = full.slice(..full.len() - 20);
assert!(handle_cb_compound(&mut truncated, &session_id, &tx, &slots).is_none());
assert!(rx.try_recv().is_err());
}
}