use std::fmt;
use log::debug;
use crate::client::connection::Connection;
use crate::client::tree::Tree;
use crate::client::SmbClient;
use crate::error::{Error, Result};
use crate::msg::copychunk::{
SrvCopychunk, SrvCopychunkCopy, SrvCopychunkResponse, SrvRequestResumeKeyResponse,
RESUME_KEY_LEN,
};
use crate::msg::ioctl::{
IoctlRequest, IoctlResponse, FSCTL_SRV_COPYCHUNK, FSCTL_SRV_REQUEST_RESUME_KEY,
SMB2_0_IOCTL_IS_FSCTL,
};
use crate::pack::{Pack, ReadCursor, Unpack, WriteCursor};
use crate::types::status::NtStatus;
use crate::types::{Command, FileId};
#[derive(Clone, Copy, PartialEq, Eq)]
pub struct ResumeKey([u8; RESUME_KEY_LEN]);
impl ResumeKey {
#[must_use]
pub fn from_bytes(bytes: [u8; RESUME_KEY_LEN]) -> Self {
Self(bytes)
}
#[must_use]
pub fn as_bytes(&self) -> &[u8; RESUME_KEY_LEN] {
&self.0
}
}
impl fmt::Debug for ResumeKey {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
write!(
f,
"ResumeKey({:02x}{:02x}{:02x}{:02x}..)",
self.0[0], self.0[1], self.0[2], self.0[3]
)
}
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub struct CopyChunk {
pub source_offset: u64,
pub target_offset: u64,
pub length: u32,
}
impl CopyChunk {
#[must_use]
pub fn new(source_offset: u64, target_offset: u64, length: u32) -> Self {
Self {
source_offset,
target_offset,
length,
}
}
}
impl From<CopyChunk> for SrvCopychunk {
fn from(c: CopyChunk) -> Self {
SrvCopychunk {
source_offset: c.source_offset,
target_offset: c.target_offset,
length: c.length,
}
}
}
#[derive(Clone, Copy, Debug, Default, PartialEq, Eq)]
pub struct CopyChunkResult {
pub chunks_written: u32,
pub total_bytes_written: u64,
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub struct ServerSideCopyLimits {
pub max_chunks: u32,
pub max_chunk_size: u32,
pub max_data_size: u32,
}
impl ServerSideCopyLimits {
pub const CONSERVATIVE: Self = Self {
max_chunks: 16,
max_chunk_size: 1024 * 1024,
max_data_size: 16 * 1024 * 1024,
};
fn sanitized(self) -> Self {
let max_chunk_size = self.max_chunk_size.max(1);
Self {
max_chunks: self.max_chunks.max(1),
max_chunk_size,
max_data_size: self.max_data_size.max(max_chunk_size),
}
}
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub enum CopyChunkOutcome {
Copied(CopyChunkResult),
Rejected {
limits: ServerSideCopyLimits,
},
}
impl Tree {
pub async fn request_resume_key(
&self,
conn: &mut Connection,
source: FileId,
) -> Result<ResumeKey> {
let req = IoctlRequest {
ctl_code: FSCTL_SRV_REQUEST_RESUME_KEY,
file_id: source,
max_input_response: 0,
max_output_response: 32,
flags: SMB2_0_IOCTL_IS_FSCTL,
input_data: Vec::new(),
};
let frame = conn
.execute(Command::Ioctl, &req, Some(self.tree_id))
.await?;
if frame.header.status != NtStatus::SUCCESS {
return Err(Error::Protocol {
status: frame.header.status,
command: Command::Ioctl,
});
}
let mut cursor = ReadCursor::new(&frame.body);
let ioctl_resp = IoctlResponse::unpack(&mut cursor)?;
let mut key_cursor = ReadCursor::new(&ioctl_resp.output_data);
let resume = SrvRequestResumeKeyResponse::unpack(&mut key_cursor)?;
Ok(ResumeKey(resume.resume_key))
}
pub async fn copy_chunks(
&self,
conn: &mut Connection,
dest: FileId,
source_key: &ResumeKey,
chunks: &[CopyChunk],
) -> Result<CopyChunkOutcome> {
let copy = SrvCopychunkCopy {
source_key: source_key.0,
chunks: chunks.iter().copied().map(SrvCopychunk::from).collect(),
};
let mut w = WriteCursor::new();
copy.pack(&mut w);
let input_data = w.into_inner();
let req = IoctlRequest {
ctl_code: FSCTL_SRV_COPYCHUNK,
file_id: dest,
max_input_response: 0,
max_output_response: SrvCopychunkResponse::SIZE as u32,
flags: SMB2_0_IOCTL_IS_FSCTL,
input_data,
};
let frame = conn
.execute(Command::Ioctl, &req, Some(self.tree_id))
.await?;
let status = frame.header.status;
if status == NtStatus::INVALID_PARAMETER {
if let Ok(resp) = parse_copychunk_response(&frame.body) {
let limits = ServerSideCopyLimits {
max_chunks: resp.chunks_written,
max_chunk_size: resp.chunk_bytes_written,
max_data_size: resp.total_bytes_written,
};
return Ok(CopyChunkOutcome::Rejected { limits });
}
return Err(Error::Protocol {
status,
command: Command::Ioctl,
});
}
if status != NtStatus::SUCCESS {
return Err(Error::Protocol {
status,
command: Command::Ioctl,
});
}
let resp = parse_copychunk_response(&frame.body)?;
Ok(CopyChunkOutcome::Copied(CopyChunkResult {
chunks_written: resp.chunks_written,
total_bytes_written: u64::from(resp.total_bytes_written),
}))
}
pub async fn server_side_copy_range(
&self,
conn: &mut Connection,
dest: FileId,
source_key: &ResumeKey,
source_offset: u64,
dest_offset: u64,
length: u64,
) -> Result<u64> {
if length == 0 {
return Ok(0);
}
let mut limits = ServerSideCopyLimits::CONSERVATIVE.sanitized();
let mut copied: u64 = 0;
while copied < length {
let remaining = length - copied;
let base_src = source_offset + copied;
let base_dst = dest_offset + copied;
let (chunks, batch_len) = build_batch(base_src, base_dst, remaining, limits);
debug_assert!(!chunks.is_empty() && batch_len > 0);
match self.copy_chunks(conn, dest, source_key, &chunks).await? {
CopyChunkOutcome::Copied(result) => {
if result.total_bytes_written == 0 {
return Err(Error::invalid_data(
"server-side copy reported success but wrote 0 bytes",
));
}
copied += result.total_bytes_written.min(batch_len);
}
CopyChunkOutcome::Rejected { limits: advertised } => {
let sane = advertised.sanitized();
if sane == limits {
return Err(Error::invalid_data(
"server rejected server-side copy without advertising smaller limits",
));
}
debug!(
"copy: server-side copy limits renegotiated to \
max_chunks={} max_chunk_size={} max_data_size={}",
sane.max_chunks, sane.max_chunk_size, sane.max_data_size
);
limits = sane;
}
}
}
Ok(copied)
}
pub async fn server_side_copy_file(
&self,
conn: &mut Connection,
source_path: &str,
dest_path: &str,
) -> Result<u64> {
self.copy_paths(conn, source_path, dest_path, true, 0, 0, None)
.await
}
pub async fn server_side_copy_file_range(
&self,
conn: &mut Connection,
source_path: &str,
source_offset: u64,
dest_path: &str,
dest_offset: u64,
length: u64,
) -> Result<u64> {
self.copy_paths(
conn,
source_path,
dest_path,
false,
source_offset,
dest_offset,
Some(length),
)
.await
}
async fn copy_paths(
&self,
conn: &mut Connection,
source_path: &str,
dest_path: &str,
truncate_dest: bool,
source_offset: u64,
dest_offset: u64,
length: Option<u64>,
) -> Result<u64> {
let src_norm = self.format_path(source_path);
let dst_norm = self.format_path(dest_path);
debug!(
"copy: server-side copy {} -> {} (truncate_dest={})",
src_norm, dst_norm, truncate_dest
);
let (src_id, src_size) = self.open_file(conn, &src_norm).await?;
let dst_open = if truncate_dest {
self.open_readwrite_overwrite(conn, &dst_norm).await
} else {
self.open_file_readwrite(conn, &dst_norm).await
};
let dst_id = match dst_open {
Ok((id, _)) => id,
Err(e) => {
let _ = self.close_handle(conn, src_id).await;
return Err(e);
}
};
let len = length.unwrap_or(src_size);
let result = async {
let key = self.request_resume_key(conn, src_id).await?;
let copied = self
.server_side_copy_range(conn, dst_id, &key, source_offset, dest_offset, len)
.await?;
self.flush_handle(conn, dst_id).await?;
Ok::<u64, Error>(copied)
}
.await;
let _ = self.close_handle(conn, dst_id).await;
let _ = self.close_handle(conn, src_id).await;
result
}
}
impl SmbClient {
pub async fn server_side_copy_file(
&mut self,
tree: &Tree,
source_path: &str,
dest_path: &str,
) -> Result<u64> {
let t = tree.clone();
let conn = self.connection_for_tree(&t);
t.server_side_copy_file(conn, source_path, dest_path).await
}
pub async fn server_side_copy_file_range(
&mut self,
tree: &Tree,
source_path: &str,
source_offset: u64,
dest_path: &str,
dest_offset: u64,
length: u64,
) -> Result<u64> {
let t = tree.clone();
let conn = self.connection_for_tree(&t);
t.server_side_copy_file_range(
conn,
source_path,
source_offset,
dest_path,
dest_offset,
length,
)
.await
}
}
fn parse_copychunk_response(frame_body: &[u8]) -> Result<SrvCopychunkResponse> {
let mut cursor = ReadCursor::new(frame_body);
let ioctl_resp = IoctlResponse::unpack(&mut cursor)?;
let mut out = ReadCursor::new(&ioctl_resp.output_data);
SrvCopychunkResponse::unpack(&mut out)
}
fn build_batch(
base_src: u64,
base_dst: u64,
remaining: u64,
limits: ServerSideCopyLimits,
) -> (Vec<CopyChunk>, u64) {
let max_data = u64::from(limits.max_data_size);
let max_chunk = u64::from(limits.max_chunk_size);
let mut chunks = Vec::new();
let mut off: u64 = 0;
while (chunks.len() as u32) < limits.max_chunks && off < remaining && off < max_data {
let this = max_chunk.min(remaining - off).min(max_data - off);
if this == 0 {
break;
}
chunks.push(CopyChunk {
source_offset: base_src + off,
target_offset: base_dst + off,
length: this as u32,
});
off += this;
}
(chunks, off)
}
#[cfg(test)]
mod tests {
use super::*;
use crate::client::test_helpers::{
build_close_response, build_create_response, build_flush_response,
build_ioctl_error_response, build_ioctl_response, build_ioctl_response_status,
setup_connection,
};
use crate::msg::copychunk::SrvCopychunkResponse;
use crate::pack::{Pack, WriteCursor};
use crate::transport::MockTransport;
use crate::types::status::NtStatus;
use crate::types::{FileId, TreeId};
use std::sync::Arc;
fn test_tree() -> Arc<Tree> {
Arc::new(Tree {
tree_id: TreeId(10),
share_name: "test".to_string(),
server: "test-server".to_string(),
is_dfs: false,
encrypt_data: false,
})
}
fn dst_id() -> FileId {
FileId {
persistent: 0xDD,
volatile: 0xEE,
}
}
fn copychunk_response_bytes(resp: SrvCopychunkResponse) -> Vec<u8> {
let mut w = WriteCursor::new();
resp.pack(&mut w);
w.into_inner()
}
fn resume_key_response_bytes(key: [u8; RESUME_KEY_LEN]) -> Vec<u8> {
let mut w = WriteCursor::new();
SrvRequestResumeKeyResponse { resume_key: key }.pack(&mut w);
w.into_inner()
}
#[tokio::test]
async fn request_resume_key_parses_key() {
let mock = Arc::new(MockTransport::new());
let key = [0x42u8; RESUME_KEY_LEN];
mock.queue_response(build_ioctl_response(
FSCTL_SRV_REQUEST_RESUME_KEY,
resume_key_response_bytes(key),
));
let mut conn = setup_connection(&mock);
let tree = test_tree();
let src = FileId {
persistent: 1,
volatile: 2,
};
let got = tree.request_resume_key(&mut conn, src).await.unwrap();
assert_eq!(got.as_bytes(), &key);
}
#[tokio::test]
async fn request_resume_key_unsupported_server_classifies_unsupported() {
let mock = Arc::new(MockTransport::new());
mock.queue_response(build_ioctl_error_response(NtStatus::NOT_SUPPORTED));
let mut conn = setup_connection(&mock);
let tree = test_tree();
let err = tree
.request_resume_key(
&mut conn,
FileId {
persistent: 1,
volatile: 2,
},
)
.await
.unwrap_err();
assert_eq!(err.kind(), crate::ErrorKind::Unsupported);
}
#[tokio::test]
async fn copy_chunks_success_reports_bytes() {
let mock = Arc::new(MockTransport::new());
mock.queue_response(build_ioctl_response(
FSCTL_SRV_COPYCHUNK,
copychunk_response_bytes(SrvCopychunkResponse {
chunks_written: 2,
chunk_bytes_written: 0,
total_bytes_written: 3000,
}),
));
let mut conn = setup_connection(&mock);
let tree = test_tree();
let key = ResumeKey::from_bytes([0; RESUME_KEY_LEN]);
let chunks = [CopyChunk::new(0, 0, 1000), CopyChunk::new(1000, 1000, 2000)];
let outcome = tree
.copy_chunks(&mut conn, dst_id(), &key, &chunks)
.await
.unwrap();
assert_eq!(
outcome,
CopyChunkOutcome::Copied(CopyChunkResult {
chunks_written: 2,
total_bytes_written: 3000,
})
);
}
#[tokio::test]
async fn copy_chunks_invalid_parameter_with_payload_is_rejected_with_limits() {
let mock = Arc::new(MockTransport::new());
mock.queue_response(build_ioctl_response_status(
FSCTL_SRV_COPYCHUNK,
NtStatus::INVALID_PARAMETER,
copychunk_response_bytes(SrvCopychunkResponse {
chunks_written: 4,
chunk_bytes_written: 65536,
total_bytes_written: 262144,
}),
));
let mut conn = setup_connection(&mock);
let tree = test_tree();
let key = ResumeKey::from_bytes([0; RESUME_KEY_LEN]);
let outcome = tree
.copy_chunks(&mut conn, dst_id(), &key, &[CopyChunk::new(0, 0, 4096)])
.await
.unwrap();
assert_eq!(
outcome,
CopyChunkOutcome::Rejected {
limits: ServerSideCopyLimits {
max_chunks: 4,
max_chunk_size: 65536,
max_data_size: 262144,
}
}
);
}
#[tokio::test]
async fn copy_chunks_invalid_parameter_without_payload_is_error() {
let mock = Arc::new(MockTransport::new());
mock.queue_response(build_ioctl_error_response(NtStatus::INVALID_PARAMETER));
let mut conn = setup_connection(&mock);
let tree = test_tree();
let key = ResumeKey::from_bytes([0; RESUME_KEY_LEN]);
let err = tree
.copy_chunks(&mut conn, dst_id(), &key, &[CopyChunk::new(0, 0, 4096)])
.await
.unwrap_err();
assert_eq!(err.status(), Some(NtStatus::INVALID_PARAMETER));
}
#[tokio::test]
async fn copy_range_batches_across_chunk_and_request_limits() {
let mib = 1024 * 1024u64;
let total = 40 * mib;
let mock = Arc::new(MockTransport::new());
for bytes in [16 * mib, 16 * mib, 8 * mib] {
mock.queue_response(build_ioctl_response(
FSCTL_SRV_COPYCHUNK,
copychunk_response_bytes(SrvCopychunkResponse {
chunks_written: (bytes / mib) as u32,
chunk_bytes_written: 0,
total_bytes_written: bytes as u32,
}),
));
}
let mut conn = setup_connection(&mock);
let tree = test_tree();
let key = ResumeKey::from_bytes([7; RESUME_KEY_LEN]);
let copied = tree
.server_side_copy_range(&mut conn, dst_id(), &key, 0, 0, total)
.await
.unwrap();
assert_eq!(copied, total);
}
#[tokio::test]
async fn copy_range_renegotiates_then_succeeds() {
let mock = Arc::new(MockTransport::new());
mock.queue_response(build_ioctl_response_status(
FSCTL_SRV_COPYCHUNK,
NtStatus::INVALID_PARAMETER,
copychunk_response_bytes(SrvCopychunkResponse {
chunks_written: 1,
chunk_bytes_written: 65536,
total_bytes_written: 65536,
}),
));
for bytes in [65536u32, 65536, 65536, 8192] {
mock.queue_response(build_ioctl_response(
FSCTL_SRV_COPYCHUNK,
copychunk_response_bytes(SrvCopychunkResponse {
chunks_written: 1,
chunk_bytes_written: 0,
total_bytes_written: bytes,
}),
));
}
let mut conn = setup_connection(&mock);
let tree = test_tree();
let key = ResumeKey::from_bytes([0; RESUME_KEY_LEN]);
let copied = tree
.server_side_copy_range(&mut conn, dst_id(), &key, 0, 0, 200 * 1024)
.await
.unwrap();
assert_eq!(copied, 200 * 1024);
}
#[tokio::test]
async fn copy_range_zero_length_is_noop() {
let mock = Arc::new(MockTransport::new());
let mut conn = setup_connection(&mock);
let tree = test_tree();
let key = ResumeKey::from_bytes([0; RESUME_KEY_LEN]);
let copied = tree
.server_side_copy_range(&mut conn, dst_id(), &key, 0, 0, 0)
.await
.unwrap();
assert_eq!(copied, 0);
}
#[tokio::test]
async fn copy_file_range_places_chunk_at_requested_offsets() {
use crate::msg::copychunk::SrvCopychunkCopy;
use crate::msg::header::Header;
use crate::msg::ioctl::{IoctlRequest, FSCTL_SRV_COPYCHUNK};
let mock = Arc::new(MockTransport::new());
let src_id = FileId {
persistent: 0xAA,
volatile: 0xBB,
};
let dst = dst_id();
let len = 1000u64;
mock.queue_response(build_create_response(src_id, 8192));
mock.queue_response(build_create_response(dst, 8192));
mock.queue_response(build_ioctl_response(
FSCTL_SRV_REQUEST_RESUME_KEY,
resume_key_response_bytes([3; RESUME_KEY_LEN]),
));
mock.queue_response(build_ioctl_response(
FSCTL_SRV_COPYCHUNK,
copychunk_response_bytes(SrvCopychunkResponse {
chunks_written: 1,
chunk_bytes_written: 0,
total_bytes_written: len as u32,
}),
));
mock.queue_response(build_flush_response());
mock.queue_response(build_close_response());
mock.queue_response(build_close_response());
let mut conn = setup_connection(&mock);
let tree = test_tree();
let copied = tree
.server_side_copy_file_range(&mut conn, "src.bin", 100, "dst.bin", 4096, len)
.await
.unwrap();
assert_eq!(copied, len);
let sent = mock.sent_messages();
let copy = sent
.iter()
.find_map(|bytes| {
let mut c = ReadCursor::new(&bytes[Header::SIZE..]);
let req = IoctlRequest::unpack(&mut c).ok()?;
if req.ctl_code != FSCTL_SRV_COPYCHUNK {
return None;
}
let mut b = ReadCursor::new(&req.input_data);
SrvCopychunkCopy::unpack(&mut b).ok()
})
.expect("a COPYCHUNK request was sent");
assert_eq!(copy.chunks.len(), 1);
assert_eq!(copy.chunks[0].source_offset, 100);
assert_eq!(copy.chunks[0].target_offset, 4096);
assert_eq!(copy.chunks[0].length, len as u32);
}
#[tokio::test]
async fn copy_file_opens_copies_and_closes_both_handles() {
let mock = Arc::new(MockTransport::new());
let src_id = FileId {
persistent: 0xAA,
volatile: 0xBB,
};
let dst = dst_id();
let size = 5000u64;
mock.queue_response(build_create_response(src_id, size));
mock.queue_response(build_create_response(dst, 0));
mock.queue_response(build_ioctl_response(
FSCTL_SRV_REQUEST_RESUME_KEY,
resume_key_response_bytes([9; RESUME_KEY_LEN]),
));
mock.queue_response(build_ioctl_response(
FSCTL_SRV_COPYCHUNK,
copychunk_response_bytes(SrvCopychunkResponse {
chunks_written: 1,
chunk_bytes_written: 0,
total_bytes_written: size as u32,
}),
));
mock.queue_response(build_flush_response());
mock.queue_response(build_close_response());
mock.queue_response(build_close_response());
let mut conn = setup_connection(&mock);
let tree = test_tree();
let copied = tree
.server_side_copy_file(&mut conn, "src.bin", "dst.bin")
.await
.unwrap();
assert_eq!(copied, size);
}
}