pub mod cancel;
pub mod devices;
pub mod io_backend;
pub mod platform;
pub mod progress;
#[cfg(not(feature = "real-io"))]
pub mod cli_simulate;
use anyhow::{Context, Result};
use log::{debug, info, warn};
use lzma::reader::LzmaReader;
use platform::{DeviceWriter, PlatformDevice};
use progress::{
check_cancel, emit_progress, OperationCancelled, OperationPhase, OperationProgress,
};
use sha2::{Digest, Sha256};
use std::fs::File;
use std::io::BufWriter;
use std::io::{BufReader, Read, Write};
use std::sync::atomic::AtomicBool;
pub fn clone<F>(
device_path: &str,
output_path: &str,
block_size: usize,
silent: bool,
mut progress: Option<F>,
cancel: Option<&AtomicBool>,
) -> Result<()>
where
F: FnMut(OperationProgress),
{
if !silent {
info!(
"Cloning device: {} to output: {} with block_size: {}",
device_path, output_path, block_size
);
}
emit_progress(
silent,
&mut progress,
OperationProgress::new(OperationPhase::Preparing)
.with_message(format!("Opening {}", device_path)),
);
let mut device_reader = PlatformDevice::new_clone_reader(device_path)?;
let total_bytes = device_reader
.device_size()
.ok()
.filter(|s| *s > 0)
.or_else(|| devices::device_size_bytes(device_path));
let output_file = File::create(output_path)
.context(format!("Failed to create output file: {}", output_path))?;
let mut writer = BufWriter::new(output_file);
let mut buffer = vec![0u8; block_size];
let mut total_bytes_read: u64 = 0;
let result = (|| -> Result<()> {
loop {
check_cancel(cancel)?;
let bytes_read = device_reader
.read(&mut buffer)
.context("Failed to read from device")?;
if bytes_read == 0 {
break;
}
writer
.write_all(&buffer[..bytes_read])
.context("Failed to write to output file")?;
total_bytes_read += bytes_read as u64;
let mut event = OperationProgress::new(OperationPhase::Writing)
.with_bytes(total_bytes_read, total_bytes);
if total_bytes.is_none() {
event = event.with_message(format!("{} bytes copied", total_bytes_read));
}
emit_progress(silent, &mut progress, event);
if !silent {
debug!("Read and written {} bytes", total_bytes_read);
}
}
Ok(())
})();
if let Err(error) = result {
drop(writer);
if error.downcast_ref::<OperationCancelled>().is_some() {
if let Err(remove_error) = std::fs::remove_file(output_path) {
warn!(
"Failed to remove incomplete clone output {}: {}",
output_path, remove_error
);
}
}
return Err(error);
}
writer
.flush()
.context("Failed to flush clone output file")?;
emit_progress(
silent,
&mut progress,
OperationProgress::new(OperationPhase::Complete)
.with_bytes(total_bytes_read, total_bytes)
.with_percentage(100.0)
.with_message("Clone completed"),
);
info!("Clone completed successfully");
Ok(())
}
pub fn flash<F>(
img_path: &str,
device_path: &str,
block_size: usize,
silent: bool,
verify: bool,
progress: Option<F>,
cancel: Option<&AtomicBool>,
) -> Result<()>
where
F: FnMut(OperationProgress),
{
if img_path.ends_with(".xz") {
info!("Detected compressed image, streaming xz flash");
flash_xz(
img_path,
device_path,
block_size,
silent,
verify,
progress,
cancel,
)
} else {
flash_image(
img_path,
device_path,
block_size,
silent,
progress,
verify,
cancel,
)
}
}
fn flash_image<F>(
img_path: &str,
device_path: &str,
block_size: usize,
silent: bool,
mut progress: Option<F>,
verify: bool,
cancel: Option<&AtomicBool>,
) -> Result<()>
where
F: FnMut(OperationProgress),
{
emit_progress(
silent,
&mut progress,
OperationProgress::new(OperationPhase::Preparing)
.with_message(format!("Opening image {}", img_path)),
);
check_cancel(cancel)?;
let img_file = File::open(img_path).context(format!("Image file not found: {}", img_path))?;
let file_size = img_file
.metadata()
.context("Failed to read image file metadata")?
.len();
let mut reader = BufReader::new(img_file);
flash_reader(
&mut reader,
device_path,
block_size,
silent,
verify,
progress,
cancel,
Some(file_size),
false,
)
}
pub fn flash_xz<F>(
img_path: &str,
device_path: &str,
block_size: usize,
silent: bool,
verify: bool,
mut progress: Option<F>,
cancel: Option<&AtomicBool>,
) -> Result<()>
where
F: FnMut(OperationProgress),
{
emit_progress(
silent,
&mut progress,
OperationProgress::new(OperationPhase::Decompressing)
.with_message(format!("Streaming decompress {}", img_path)),
);
check_cancel(cancel)?;
let input_file =
File::open(img_path).context(format!("Failed to open compressed file: {}", img_path))?;
let buffered_reader = BufReader::new(input_file);
let mut decoder = LzmaReader::new_decompressor(buffered_reader)
.context("Failed to create LZMA decompressor")?;
flash_reader(
&mut decoder,
device_path,
block_size,
silent,
verify,
progress,
cancel,
None,
true,
)?;
info!("Flash successful");
Ok(())
}
#[allow(clippy::too_many_arguments)]
fn flash_reader<R, F>(
reader: &mut R,
device_path: &str,
block_size: usize,
silent: bool,
verify: bool,
mut progress: Option<F>,
cancel: Option<&AtomicBool>,
known_size: Option<u64>,
from_compressed: bool,
) -> Result<()>
where
R: Read,
F: FnMut(OperationProgress),
{
let mut device_writer = PlatformDevice::new_writer(device_path)?;
let mut buffer = vec![0u8; block_size];
let mut source_hasher = verify.then(Sha256::new);
if !silent {
if let Some(size) = known_size {
info!("Writing image to the device... size: {}", size);
} else if from_compressed {
info!("Writing streamed decompressed image to the device...");
} else {
info!("Writing image to the device...");
}
}
let incremental_sync = device_writer.supports_incremental_sync();
let written = write_image_to_device(
reader,
device_writer.as_mut(),
&mut buffer,
known_size,
verify,
silent,
&mut progress,
cancel,
&mut source_hasher,
)?;
let file_size = known_size.unwrap_or(written);
if !incremental_sync {
let (sync_pct_start, sync_pct_end) = if verify { (85.0, 95.0) } else { (90.0, 100.0) };
let mut emit_sync_progress = |synced: u64, total: u64| {
let fraction = if total == 0 {
1.0
} else {
synced as f64 / total as f64
};
let pct = sync_pct_start + fraction * (sync_pct_end - sync_pct_start);
emit_progress(
silent,
&mut progress,
OperationProgress::new(OperationPhase::Syncing)
.with_bytes(synced, Some(total))
.with_percentage(pct)
.with_message("Syncing data to device"),
);
};
device_writer
.flush_and_sync_with_progress(file_size, Some(&mut emit_sync_progress))
.context("Failed to flush and sync device")?;
}
if !verify {
emit_progress(
silent,
&mut progress,
OperationProgress::new(OperationPhase::Complete)
.with_bytes(file_size, Some(file_size))
.with_percentage(100.0)
.with_message("Flash completed"),
);
if !silent {
info!("Flash completed successfully");
}
return Ok(());
}
let img_checksum = source_hasher
.map(|hasher| format!("{:x}", hasher.finalize()))
.context("Missing source checksum")?;
if !silent {
info!("Source image checksum: {}", img_checksum);
}
emit_progress(
silent,
&mut progress,
OperationProgress::new(OperationPhase::Verifying)
.with_percentage(95.0)
.with_message("Verifying checksum"),
);
let mut verified: u64 = 0;
let verify_hasher = if device_writer.supports_inline_verify() {
device_writer
.rewind_for_verify()
.context("Failed to prepare device for verification")?;
let mut inline_reader = InlineVerifyReader {
writer: device_writer.as_mut(),
};
verify_checksum_with_progress(
&mut inline_reader,
file_size,
silent,
&mut progress,
&mut verified,
cancel,
)?
} else {
let device_reader = PlatformDevice::new_verify_reader(device_path)?;
let mut buffered_reader = BufReader::with_capacity(1024 * 1024, device_reader);
verify_checksum_with_progress(
&mut buffered_reader,
file_size,
silent,
&mut progress,
&mut verified,
cancel,
)?
};
let device_checksum = format!("{:x}", verify_hasher.finalize());
if !silent {
info!("Device checksum: {}", device_checksum);
}
if img_checksum == device_checksum {
emit_progress(
silent,
&mut progress,
OperationProgress::new(OperationPhase::Complete)
.with_percentage(100.0)
.with_message("Checksums match"),
);
if !silent {
info!("Checksums match. Write operation successful.");
}
Ok(())
} else {
emit_progress(
silent,
&mut progress,
OperationProgress::new(OperationPhase::Failed).with_message("Checksums do not match"),
);
log::error!("Checksums do not match. Write operation may have failed.");
anyhow::bail!("Checksums do not match");
}
}
fn verify_checksum_with_progress<F>(
reader: &mut dyn Read,
size: u64,
silent: bool,
progress: &mut Option<F>,
verified: &mut u64,
cancel: Option<&AtomicBool>,
) -> Result<Sha256>
where
F: FnMut(OperationProgress),
{
let mut hasher = Sha256::new();
let mut buffer = vec![0u8; 65536];
let mut remaining = size;
while remaining > 0 {
check_cancel(cancel)?;
let to_read = usize::try_from(remaining.min(buffer.len() as u64))
.context("Verify read chunk too large")?;
let bytes_read = reader.read(&mut buffer[..to_read]).with_context(|| {
format!(
"Failed to read from device during verification ({} bytes remaining)",
remaining
)
})?;
if bytes_read == 0 {
if remaining > 0 {
anyhow::bail!(
"Unexpected end of device read during verification ({} bytes short)",
remaining
);
}
break;
}
hasher.update(&buffer[..bytes_read]);
remaining -= bytes_read as u64;
*verified += bytes_read as u64;
let verify_pct = if size == 0 {
99.9
} else {
95.0 + (*verified as f64 / size as f64) * 5.0
};
emit_progress(
silent,
progress,
OperationProgress::new(OperationPhase::Verifying)
.with_bytes(*verified, Some(size))
.with_percentage(verify_pct.min(99.9)),
);
}
Ok(hasher)
}
const INCREMENTAL_SYNC_INTERVAL: u64 = 4 * 1024 * 1024;
fn sync_pending_device_bytes(
device_writer: &mut dyn DeviceWriter,
sync_cursor: &mut u64,
device_written_through: u64,
image_bytes_durable: &mut u64,
) -> Result<()> {
while device_written_through.saturating_sub(*sync_cursor) >= INCREMENTAL_SYNC_INTERVAL {
device_writer
.sync_written_range(*sync_cursor, INCREMENTAL_SYNC_INTERVAL)
.context("Failed to sync written device range")?;
*sync_cursor += INCREMENTAL_SYNC_INTERVAL;
*image_bytes_durable += INCREMENTAL_SYNC_INTERVAL;
}
Ok(())
}
fn flash_progress_percentage(
durable_bytes: u64,
file_size: Option<u64>,
verify: bool,
incremental_sync: bool,
) -> Option<f64> {
let file_size = file_size.filter(|s| *s > 0)?;
let fraction = durable_bytes as f64 / file_size as f64;
let pct_end = if incremental_sync {
if verify { 95.0 } else { 100.0 }
} else if verify {
85.0
} else {
90.0
};
Some(fraction * pct_end)
}
#[allow(clippy::too_many_arguments)]
fn write_image_to_device<R, F>(
reader: &mut R,
device_writer: &mut dyn DeviceWriter,
buffer: &mut [u8],
known_size: Option<u64>,
verify: bool,
silent: bool,
progress: &mut Option<F>,
cancel: Option<&AtomicBool>,
source_hasher: &mut Option<Sha256>,
) -> Result<u64>
where
R: Read,
F: FnMut(OperationProgress),
{
let defer_partition_table = known_size.map(|s| s > 0).unwrap_or(true);
let incremental_sync = device_writer.supports_incremental_sync();
let mut count: u64 = 0;
let mut device_offset: u64 = 0;
let mut deferred_header: Option<Vec<u8>> = None;
let mut sync_cursor: u64 = 0;
let mut image_bytes_durable: u64 = 0;
if defer_partition_table {
check_cancel(cancel)?;
let bytes_read = reader.read(buffer).context("Failed to read image file")?;
if bytes_read == 0 {
return Ok(0);
}
deferred_header = Some(buffer[..bytes_read].to_vec());
device_offset = bytes_read as u64;
count += bytes_read as u64;
if incremental_sync {
sync_cursor = device_offset;
}
if let Some(hasher) = source_hasher.as_mut() {
hasher.update(&buffer[..bytes_read]);
}
if !silent {
info!("Deferring first {bytes_read} bytes (partition table) until end of write");
}
}
loop {
check_cancel(cancel)?;
let bytes_read = reader.read(buffer).context("Failed to read image file")?;
if bytes_read == 0 {
break;
}
if let Some(hasher) = source_hasher.as_mut() {
hasher.update(&buffer[..bytes_read]);
}
if defer_partition_table {
device_writer.write_at(device_offset, &buffer[..bytes_read])?
} else if let Err(err) = device_writer.write_all(&buffer[..bytes_read]) {
let message = format_write_device_error(&err, count, bytes_read);
return Err(anyhow::Error::new(err).context(message));
}
if defer_partition_table {
device_offset += bytes_read as u64;
}
count += bytes_read as u64;
let device_written_through = if defer_partition_table {
device_offset
} else {
count
};
if incremental_sync {
sync_pending_device_bytes(
device_writer,
&mut sync_cursor,
device_written_through,
&mut image_bytes_durable,
)?;
}
let progress_bytes = if incremental_sync {
image_bytes_durable
} else {
count
};
let mut event = OperationProgress::new(OperationPhase::Writing)
.with_bytes(progress_bytes, known_size);
if let Some(pct) =
flash_progress_percentage(progress_bytes, known_size, verify, incremental_sync)
{
event = event.with_percentage(pct);
} else {
event = event.with_message(format!("{progress_bytes} bytes written"));
}
emit_progress(silent, progress, event);
if !silent {
debug!(
"Written {} durable {} known_size {:?}",
count, image_bytes_durable, known_size
);
}
}
let device_written_through = if defer_partition_table {
device_offset
} else {
count
};
if let Some(header) = deferred_header {
check_cancel(cancel)?;
let header_len = header.len();
if let Err(err) = device_writer.write_at(0, &header) {
return Err(err.context(format!(
"Failed to write deferred partition table at offset 0 ({header_len} bytes)"
)));
}
if !silent {
info!("Wrote deferred partition table ({header_len} bytes) at offset 0");
}
if incremental_sync {
let header_len_u64 = u64::try_from(header_len).context("Partition table too large")?;
device_writer
.sync_written_range(0, header_len_u64)
.context("Failed to sync deferred partition table")?;
image_bytes_durable += header_len_u64;
sync_cursor = sync_cursor.max(header_len_u64);
}
}
if incremental_sync {
if device_written_through > sync_cursor {
let tail_len = device_written_through - sync_cursor;
device_writer
.sync_written_range(sync_cursor, tail_len)
.context("Failed to sync remaining device bytes")?;
image_bytes_durable += tail_len;
}
device_writer
.flush_and_sync()
.context("Failed to finalize device sync")?;
let mut event = OperationProgress::new(OperationPhase::Writing)
.with_bytes(image_bytes_durable, known_size.or(Some(count)));
if let Some(pct) =
flash_progress_percentage(image_bytes_durable, known_size, verify, true)
{
event = event.with_percentage(pct);
}
emit_progress(silent, progress, event);
}
Ok(count)
}
struct InlineVerifyReader<'a> {
writer: &'a mut dyn DeviceWriter,
}
impl Read for InlineVerifyReader<'_> {
fn read(&mut self, buf: &mut [u8]) -> std::io::Result<usize> {
self.writer.read_for_verify(buf)
}
}
fn format_write_device_error(err: &std::io::Error, offset: u64, nbytes: usize) -> String {
let mut message =
format!("Failed to write to device at byte offset {offset} ({nbytes} bytes): {err}");
if let Some(hint) = devices::hint_for_io_error(err) {
message.push_str(&hint);
}
message
}