use crate::upload::blobstore::Error;
use async_trait::async_trait;
use azure_core::{
Bytes,
http::{NoFormat, RequestContent, XmlFormat},
};
use azure_storage_blob::{
BlockBlobClient,
models::{
BlockBlobClientCommitBlockListOptions, BlockBlobClientStageBlockOptions, BlockLookupList,
},
};
use core::num::NonZeroUsize;
use core::{
pin::Pin,
task::{Context, Poll},
};
use std::{
io::{Result as IoResult, Write},
sync::{Arc, Mutex},
};
use tokio::{
io::AsyncWrite,
runtime::Handle,
sync::{Semaphore, mpsc},
task::JoinHandle,
};
use tokio_util::io::SyncIoBridge;
type Result<T> = core::result::Result<T, Error>;
fn block_id(index: u64) -> Vec<u8> {
index.to_be_bytes().to_vec()
}
#[expect(
clippy::redundant_pub_crate,
reason = "appears in the signature of pub(crate) BlockBlobStream::with_stager"
)]
#[async_trait]
pub(crate) trait BlockStager: Send + Sync + 'static {
async fn stage_block(&self, block_id: Vec<u8>, body: Bytes) -> Result<()>;
async fn commit_block_list(&self, block_ids: Vec<Vec<u8>>) -> Result<()>;
}
struct SdkStager {
client: Arc<BlockBlobClient>,
}
#[async_trait]
impl BlockStager for SdkStager {
async fn stage_block(&self, block_id: Vec<u8>, body: Bytes) -> Result<()> {
let len = u64::try_from(body.len())?;
let content: RequestContent<Bytes, NoFormat> = body.into();
self.client
.stage_block(
&block_id,
len,
content,
Option::<BlockBlobClientStageBlockOptions<'_>>::None,
)
.await?;
Ok(())
}
async fn commit_block_list(&self, block_ids: Vec<Vec<u8>>) -> Result<()> {
let list = BlockLookupList {
latest: Some(block_ids),
..Default::default()
};
let content: RequestContent<BlockLookupList, XmlFormat> = list.try_into()?;
self.client
.commit_block_list(
content,
Option::<BlockBlobClientCommitBlockListOptions<'_>>::None,
)
.await?;
Ok(())
}
}
enum UploaderMsg {
Stage { index: u64, data: Bytes },
}
struct UploaderResult {
completed: Vec<u64>,
first_error: Option<Error>,
}
type ReservationFuture = Pin<
Box<
dyn Future<
Output = core::result::Result<
mpsc::OwnedPermit<UploaderMsg>,
mpsc::error::SendError<()>,
>,
> + Send,
>,
>;
struct BlockBlobAsyncWriter {
sender: Option<mpsc::Sender<UploaderMsg>>,
buf: Vec<u8>,
block_size: usize,
max_blocks: u64,
next_index: u64,
error_slot: Arc<Mutex<Option<Error>>>,
pending_reservation: Option<ReservationFuture>,
}
impl BlockBlobAsyncWriter {
fn new(
sender: mpsc::Sender<UploaderMsg>,
block_size: NonZeroUsize,
max_blocks: u64,
error_slot: Arc<Mutex<Option<Error>>>,
) -> Self {
Self {
sender: Some(sender),
buf: Vec::with_capacity(block_size.get()),
block_size: block_size.get(),
max_blocks,
next_index: 0,
error_slot,
pending_reservation: None,
}
}
fn first_error_io(&self) -> Option<std::io::Error> {
if let Ok(slot) = self.error_slot.lock() {
slot.as_ref()
.map(ToString::to_string)
.map(std::io::Error::other)
} else {
None
}
}
fn record_error_and_close(&mut self, err: Error) -> std::io::Error {
let message = err.to_string();
if let Ok(mut slot) = self.error_slot.lock()
&& slot.is_none()
{
*slot = Some(err);
}
self.sender = None;
std::io::Error::other(message)
}
fn allocate_index(&mut self) -> Result<u64> {
if self.next_index >= self.max_blocks {
return Err(Error::TooLarge);
}
let index = self.next_index;
self.next_index = self.next_index.checked_add(1).ok_or(Error::TooLarge)?;
Ok(index)
}
fn try_dispatch(&mut self, cx: &mut Context<'_>) -> Poll<IoResult<()>> {
if self.buf.len() < self.block_size {
return Poll::Ready(Ok(()));
}
if self.pending_reservation.is_none() {
let Some(sender) = self.sender.as_ref() else {
return Poll::Ready(Err(std::io::Error::other(
"blob writer was already shut down",
)));
};
let sender = sender.clone();
self.pending_reservation = Some(Box::pin(sender.reserve_owned()));
}
let mut reservation = self
.pending_reservation
.take()
.ok_or_else(|| std::io::Error::other("missing reservation slot"))?;
match reservation.as_mut().poll(cx) {
Poll::Pending => {
self.pending_reservation = Some(reservation);
Poll::Pending
}
Poll::Ready(Err(_send)) => {
self.sender = None;
let err = self
.first_error_io()
.unwrap_or_else(|| std::io::Error::other("uploader exited early"));
Poll::Ready(Err(err))
}
Poll::Ready(Ok(permit)) => {
let index = match self.allocate_index() {
Ok(index) => index,
Err(err) => return Poll::Ready(Err(self.record_error_and_close(err))),
};
let block_size = self.block_size;
let data = core::mem::replace(&mut self.buf, Vec::with_capacity(block_size));
permit.send(UploaderMsg::Stage {
index,
data: Bytes::from(data),
});
Poll::Ready(Ok(()))
}
}
}
}
impl AsyncWrite for BlockBlobAsyncWriter {
fn poll_write(
mut self: Pin<&mut Self>,
cx: &mut Context<'_>,
buf: &[u8],
) -> Poll<IoResult<usize>> {
if let Some(err) = self.first_error_io() {
return Poll::Ready(Err(err));
}
if self.buf.len() >= self.block_size {
match self.try_dispatch(cx) {
Poll::Pending => return Poll::Pending,
Poll::Ready(Err(e)) => return Poll::Ready(Err(e)),
Poll::Ready(Ok(())) => {}
}
}
let take = self
.block_size
.saturating_sub(self.buf.len())
.min(buf.len());
self.buf.extend_from_slice(buf.get(..take).unwrap_or(&[]));
Poll::Ready(Ok(take))
}
fn poll_flush(self: Pin<&mut Self>, _cx: &mut Context<'_>) -> Poll<IoResult<()>> {
Poll::Ready(Ok(()))
}
fn poll_shutdown(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<IoResult<()>> {
if let Some(err) = self.first_error_io() {
return Poll::Ready(Err(err));
}
if !self.buf.is_empty() {
if self.pending_reservation.is_none() {
let Some(sender) = self.sender.as_ref() else {
return Poll::Ready(Ok(()));
};
let sender = sender.clone();
self.pending_reservation = Some(Box::pin(sender.reserve_owned()));
}
let mut reservation = self
.pending_reservation
.take()
.ok_or_else(|| std::io::Error::other("missing reservation slot"))?;
match reservation.as_mut().poll(cx) {
Poll::Pending => {
self.pending_reservation = Some(reservation);
return Poll::Pending;
}
Poll::Ready(Err(_)) => {
self.sender = None;
let err = self
.first_error_io()
.unwrap_or_else(|| std::io::Error::other("uploader exited early"));
return Poll::Ready(Err(err));
}
Poll::Ready(Ok(permit)) => {
let index = match self.allocate_index() {
Ok(index) => index,
Err(err) => return Poll::Ready(Err(self.record_error_and_close(err))),
};
let block_size = self.block_size;
let data = core::mem::replace(&mut self.buf, Vec::with_capacity(block_size));
permit.send(UploaderMsg::Stage {
index,
data: Bytes::from(data),
});
}
}
}
self.sender = None;
Poll::Ready(Ok(()))
}
}
pub struct BlockBlobStream {
bridge: SyncIoBridge<BlockBlobAsyncWriter>,
uploader: Option<JoinHandle<UploaderResult>>,
stager: Arc<dyn BlockStager>,
}
pub const BLOB_MAX_BLOCKS: u64 = 50_000;
impl BlockBlobStream {
#[must_use]
pub fn new(
client: BlockBlobClient,
block_size: NonZeroUsize,
concurrency: NonZeroUsize,
) -> Self {
Self::with_stager(
Arc::new(SdkStager {
client: Arc::new(client),
}),
block_size,
concurrency,
)
}
pub(crate) fn with_stager(
stager: Arc<dyn BlockStager>,
block_size: NonZeroUsize,
concurrency: NonZeroUsize,
) -> Self {
Self::with_stager_and_max_blocks(stager, block_size, concurrency, BLOB_MAX_BLOCKS)
}
pub(crate) fn with_stager_and_max_blocks(
stager: Arc<dyn BlockStager>,
block_size: NonZeroUsize,
concurrency: NonZeroUsize,
max_blocks: u64,
) -> Self {
let handle = Handle::current();
let error_slot = Arc::new(Mutex::new(None));
let (tx, rx) = mpsc::channel::<UploaderMsg>(concurrency.get());
let uploader = handle.spawn(run_uploader(
stager.clone(),
rx,
Arc::new(Semaphore::new(concurrency.get())),
error_slot.clone(),
));
let writer = BlockBlobAsyncWriter::new(tx, block_size, max_blocks, error_slot);
let bridge = SyncIoBridge::new_with_handle(writer, handle);
Self {
bridge,
uploader: Some(uploader),
stager,
}
}
pub fn writer(&mut self) -> &mut dyn Write {
&mut self.bridge
}
pub fn finish_writes(&mut self) -> IoResult<()> {
self.bridge.shutdown()
}
pub async fn finalize(mut self) -> Result<()> {
let result = self.await_uploader().await?;
if let Some(err) = result.first_error {
return Err(err);
}
let mut indices = result.completed;
indices.sort_unstable();
let block_ids: Vec<Vec<u8>> = indices.into_iter().map(block_id).collect();
self.stager.commit_block_list(block_ids).await
}
pub async fn abort(self) -> Result<()> {
let result = Self::await_uploader_handle(self.close_for_abort()).await?;
if let Some(err) = result.first_error {
return Err(err);
}
Ok(())
}
async fn await_uploader(&mut self) -> Result<UploaderResult> {
Self::await_uploader_handle(self.uploader.take()).await
}
fn close_for_abort(self) -> Option<JoinHandle<UploaderResult>> {
let Self {
bridge: _closed_bridge,
uploader,
stager: _,
} = self;
uploader
}
async fn await_uploader_handle(
uploader: Option<JoinHandle<UploaderResult>>,
) -> Result<UploaderResult> {
let Some(uploader) = uploader else {
return Ok(UploaderResult {
completed: Vec::new(),
first_error: None,
});
};
uploader
.await
.map_err(|e| Error::Io(std::io::Error::other(e.to_string())))
}
}
async fn run_uploader(
stager: Arc<dyn BlockStager>,
mut rx: mpsc::Receiver<UploaderMsg>,
semaphore: Arc<Semaphore>,
error_slot: Arc<Mutex<Option<Error>>>,
) -> UploaderResult {
let mut in_flight: Vec<JoinHandle<core::result::Result<u64, (u64, Error)>>> = Vec::new();
while let Some(msg) = rx.recv().await {
match msg {
UploaderMsg::Stage { index, data } => {
let Ok(permit) = semaphore.clone().acquire_owned().await else {
break;
};
let stager = stager.clone();
let id = block_id(index);
let worker = tokio::spawn(async move {
let _permit = permit;
stager
.stage_block(id, data)
.await
.map(|()| index)
.map_err(|e| (index, e))
});
in_flight.push(worker);
}
}
}
let mut completed = Vec::with_capacity(in_flight.len());
for handle in in_flight {
match handle.await {
Ok(Ok(index)) => completed.push(index),
Ok(Err((_index, err))) => {
if let Ok(mut slot) = error_slot.lock()
&& slot.is_none()
{
*slot = Some(err);
}
}
Err(join_err) => {
if let Ok(mut slot) = error_slot.lock()
&& slot.is_none()
{
*slot = Some(Error::Io(std::io::Error::other(join_err.to_string())));
}
}
}
}
let first_error = error_slot.lock().ok().and_then(|mut s| s.take());
UploaderResult {
completed,
first_error,
}
}
#[cfg(test)]
mod tests {
#![expect(
clippy::expect_used,
clippy::indexing_slicing,
clippy::similar_names,
reason = "tests assert on pre-known shapes and value counts"
)]
use super::*;
use crate::{
image::{Format, Header, Image, MAX_BLOCK_SIZE},
snapshot::{Snapshot, Source},
};
use core::{
ops::Range,
sync::atomic::{AtomicUsize, Ordering},
};
use std::{
fs,
io::{Cursor, Read as _, Seek as _, SeekFrom},
path::Path,
sync::Mutex as StdMutex,
};
struct FakeStager {
staged: StdMutex<Vec<(Vec<u8>, Bytes)>>,
commits: StdMutex<Vec<Vec<Vec<u8>>>>,
fail_index: Option<u64>,
stage_call_count: AtomicUsize,
max_concurrent_stages: AtomicUsize,
current_stages: AtomicUsize,
}
impl FakeStager {
fn new() -> Self {
Self {
staged: StdMutex::new(Vec::new()),
commits: StdMutex::new(Vec::new()),
fail_index: None,
stage_call_count: AtomicUsize::new(0),
max_concurrent_stages: AtomicUsize::new(0),
current_stages: AtomicUsize::new(0),
}
}
fn failing(index: u64) -> Self {
Self {
fail_index: Some(index),
..Self::new()
}
}
fn locked_staged(&self) -> Vec<(Vec<u8>, Bytes)> {
self.staged
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner)
.clone()
}
fn locked_commits(&self) -> Vec<Vec<Vec<u8>>> {
self.commits
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner)
.clone()
}
}
#[async_trait]
impl BlockStager for FakeStager {
async fn stage_block(&self, block_id: Vec<u8>, body: Bytes) -> Result<()> {
self.stage_call_count.fetch_add(1, Ordering::SeqCst);
let now = self
.current_stages
.fetch_add(1, Ordering::SeqCst)
.saturating_add(1);
self.max_concurrent_stages.fetch_max(now, Ordering::SeqCst);
tokio::task::yield_now().await;
tokio::task::yield_now().await;
let should_fail = self.fail_index.is_some_and(|target| {
if block_id.len() == 8 {
let mut buf = [0_u8; 8];
buf.copy_from_slice(&block_id);
u64::from_be_bytes(buf) == target
} else {
false
}
});
self.current_stages.fetch_sub(1, Ordering::SeqCst);
if should_fail {
return Err(Error::Io(std::io::Error::other("simulated failure")));
}
self.staged
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner)
.push((block_id, body));
Ok(())
}
async fn commit_block_list(&self, block_ids: Vec<Vec<u8>>) -> Result<()> {
self.commits
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner)
.push(block_ids);
Ok(())
}
}
fn nz(n: usize) -> NonZeroUsize {
NonZeroUsize::new(n).expect("test constant non-zero")
}
fn build_stream(
stager: Arc<FakeStager>,
block_size: usize,
concurrency: usize,
) -> BlockBlobStream {
BlockBlobStream::with_stager(stager, nz(block_size), nz(concurrency))
}
fn build_stream_with_max_blocks(
stager: Arc<FakeStager>,
block_size: usize,
concurrency: usize,
max_blocks: u64,
) -> BlockBlobStream {
BlockBlobStream::with_stager_and_max_blocks(
stager,
nz(block_size),
nz(concurrency),
max_blocks,
)
}
async fn run_write<F>(stream: BlockBlobStream, write: F) -> (BlockBlobStream, IoResult<()>)
where
F: FnOnce(&mut dyn Write) -> IoResult<()> + Send + 'static,
{
tokio::task::spawn_blocking(move || {
let mut stream = stream;
let result = write(stream.writer());
let shutdown = stream.finish_writes();
let combined = result.and(shutdown);
(stream, combined)
})
.await
.expect("spawn_blocking join")
}
struct SnapshotFixture {
_dir: tempfile::TempDir,
source_path: std::path::PathBuf,
payload: Vec<u8>,
memory_ranges: Vec<Range<u64>>,
expected_ranges: Vec<Range<u64>>,
}
impl SnapshotFixture {
fn new() -> Self {
let max_block_size =
usize::try_from(MAX_BLOCK_SIZE).expect("MAX_BLOCK_SIZE fits usize");
let first_range_end = max_block_size
.checked_add(8_192)
.expect("test size fits usize");
let zero_range_end = first_range_end
.checked_add(8_192)
.expect("test size fits usize");
let tail_range_end = zero_range_end
.checked_add(12_345)
.expect("test size fits usize");
let mut payload = vec![0; tail_range_end];
fill_nonzero(&mut payload, 0..first_range_end);
fill_nonzero(&mut payload, zero_range_end..tail_range_end);
let dir = tempfile::TempDir::new_in(env!("CARGO_MANIFEST_DIR"))
.expect("create workspace-local temp dir");
let source_path = dir.path().join("raw-memory.bin");
fs::write(&source_path, &payload).expect("write raw memory fixture");
let max_block_size = u64::try_from(max_block_size).expect("test size fits u64");
let first_range_end = u64::try_from(first_range_end).expect("test size fits u64");
let zero_range_end = u64::try_from(zero_range_end).expect("test size fits u64");
let tail_range_end = u64::try_from(tail_range_end).expect("test size fits u64");
Self {
_dir: dir,
source_path,
payload,
memory_ranges: vec![
0..first_range_end,
first_range_end..zero_range_end,
zero_range_end..tail_range_end,
],
expected_ranges: vec![
0..max_block_size,
max_block_size..first_range_end,
zero_range_end..tail_range_end,
],
}
}
}
fn fill_nonzero(payload: &mut [u8], range: Range<usize>) {
let slice = payload
.get_mut(range)
.expect("test range lies within payload");
for (idx, byte) in slice.iter_mut().enumerate() {
let value = u8::try_from(idx % 251)
.expect("modulo constrains byte")
.saturating_add(1);
*byte = value;
}
}
fn committed_stream(stager: &FakeStager, blob_block_size: usize) -> Vec<u8> {
let staged = stager.locked_staged();
let commits = stager.locked_commits();
assert_eq!(commits.len(), 1, "commit called exactly once");
let committed_ids = &commits[0];
let mut sorted_committed_ids = committed_ids.clone();
sorted_committed_ids.sort();
assert_eq!(
committed_ids, &sorted_committed_ids,
"committed block ids are sorted ascending"
);
let mut staged_ids: Vec<Vec<u8>> = staged.iter().map(|entry| entry.0.clone()).collect();
staged_ids.sort();
assert_eq!(
committed_ids, &staged_ids,
"commit includes every staged block id"
);
let mut bytes = Vec::new();
for id in committed_ids {
let entry = staged
.iter()
.find(|entry| &entry.0 == id)
.expect("committed id was staged");
bytes.extend_from_slice(&entry.1);
}
let expected_blocks = bytes.len().div_ceil(blob_block_size);
assert_eq!(
staged.len(),
expected_blocks,
"staged blocks cover exactly the committed byte stream"
);
let last_len = committed_ids
.last()
.and_then(|id| staged.iter().find(|entry| &entry.0 == id))
.map(|entry| entry.1.len())
.expect("snapshot stages at least one block");
assert!(
last_len < blob_block_size,
"trailing partial block is staged on shutdown"
);
bytes
}
fn read_headers(encoded: &[u8], expected_format: Format) -> Vec<Range<u64>> {
let mut cursor = Cursor::new(encoded);
let encoded_len = u64::try_from(encoded.len()).expect("encoded length fits u64");
let mut ranges = Vec::new();
while cursor
.stream_position()
.expect("read stream position")
.lt(&encoded_len)
{
let header = Header::read(&mut cursor).expect("snapshot header parses");
assert_eq!(header.format, expected_format);
let size = i64::try_from(header.size().expect("header size fits usize"))
.expect("header size fits i64");
ranges.push(header.range.clone());
match header.format {
Format::Lime => {
cursor
.seek(SeekFrom::Current(size))
.expect("seek past LiME payload");
}
Format::AvmlCompressed => {
let mut decoder = snap::read::FrameDecoder::new(&mut cursor)
.take(u64::try_from(size).expect("positive header size fits u64"));
std::io::copy(&mut decoder, &mut std::io::sink())
.expect("compressed payload decodes");
cursor
.seek(SeekFrom::Current(8))
.expect("seek past compressed byte count");
}
}
}
ranges
}
fn convert_to_lime(encoded: &[u8]) -> Vec<u8> {
let encoded_len = u64::try_from(encoded.len()).expect("encoded length fits u64");
let mut image =
Image::from_streams(Format::Lime, Cursor::new(encoded), Cursor::new(Vec::new()));
while image
.src
.stream_position()
.expect("read stream position")
.lt(&encoded_len)
{
image.convert_block().expect("convert block to LiME");
}
image.dst.into_inner()
}
fn assert_lime_payload_matches(lime: &[u8], expected_ranges: &[Range<u64>], payload: &[u8]) {
let mut cursor = Cursor::new(lime);
let lime_len = u64::try_from(lime.len()).expect("LiME length fits u64");
let mut actual_ranges = Vec::new();
while cursor
.stream_position()
.expect("read stream position")
.lt(&lime_len)
{
let header = Header::read(&mut cursor).expect("LiME header parses");
assert_eq!(header.format, Format::Lime);
let size = header.size().expect("header size fits usize");
let mut block = vec![0; size];
cursor.read_exact(&mut block).expect("read LiME payload");
let start = usize::try_from(header.range.start).expect("range start fits usize");
let end = usize::try_from(header.range.end).expect("range end fits usize");
assert_eq!(
block,
payload.get(start..end).expect("range lies within payload")
);
actual_ranges.push(header.range);
}
assert_eq!(actual_ranges, expected_ranges);
}
async fn assert_snapshot_streams_end_to_end(format: Format) {
let fixture = SnapshotFixture::new();
let stager = Arc::new(FakeStager::new());
let blob_block_size = 1_000_000;
let stream = build_stream(stager.clone(), blob_block_size, 4);
let source_path = fixture.source_path.clone();
let memory_ranges = fixture.memory_ranges.clone();
let (stream, result) = run_write(stream, move |writer| {
Snapshot::new(
Path::new("unused-streaming-test-destination"),
memory_ranges,
)
.source(Some(Source::Raw(source_path)))
.format(format)
.create_to_writer(writer)
.map_err(|err| std::io::Error::other(err.to_string()))
})
.await;
result.expect("snapshot writes and shuts down blob stream");
stream.finalize().await.expect("finalize commits blocks");
let reconstructed = committed_stream(&stager, blob_block_size);
assert_eq!(
read_headers(&reconstructed, format),
fixture.expected_ranges,
"headers preserve non-zero ranges and elide zero-only ranges"
);
let lime = convert_to_lime(&reconstructed);
assert_lime_payload_matches(&lime, &fixture.expected_ranges, &fixture.payload);
}
#[tokio::test(flavor = "multi_thread", worker_threads = 2)]
async fn snapshot_streams_end_to_end_through_blob_stager() {
assert_snapshot_streams_end_to_end(Format::Lime).await;
assert_snapshot_streams_end_to_end(Format::AvmlCompressed).await;
}
#[tokio::test(flavor = "multi_thread", worker_threads = 2)]
async fn finalize_with_zero_writes_commits_empty_block_list() {
let stager = Arc::new(FakeStager::new());
let stream = build_stream(stager.clone(), 8, 2);
let (stream, result) = run_write(stream, |_w| Ok(())).await;
result.expect("writer shutdown succeeds");
stream.finalize().await.expect("finalize succeeds");
let commits = stager.locked_commits();
assert_eq!(commits.len(), 1, "exactly one commit");
assert!(
commits[0].is_empty(),
"empty snapshot commits empty block list"
);
}
#[tokio::test(flavor = "multi_thread", worker_threads = 2)]
async fn rotation_at_exact_block_size_emits_uniform_ids() {
let stager = Arc::new(FakeStager::new());
let stream = build_stream(stager.clone(), 4, 3);
let payload: Vec<u8> = (0..12).collect();
let (stream, result) = run_write(stream, move |w| w.write_all(&payload)).await;
result.expect("write + shutdown");
stream.finalize().await.expect("finalize");
let staged = stager.locked_staged();
assert_eq!(staged.len(), 3, "three full blocks staged");
for entry in &staged {
assert_eq!(entry.0.len(), 8, "block ids uniform width");
}
let commits = stager.locked_commits();
assert_eq!(commits.len(), 1);
let ids = &commits[0];
let mut sorted = ids.clone();
sorted.sort();
assert_eq!(ids, &sorted, "committed ids are sorted ascending");
}
#[tokio::test(flavor = "multi_thread", worker_threads = 2)]
async fn trailing_partial_block_is_staged_on_shutdown() {
let stager = Arc::new(FakeStager::new());
let stream = build_stream(stager.clone(), 4, 3);
let payload: Vec<u8> = (0..6).collect(); let (stream, result) = run_write(stream, move |w| w.write_all(&payload)).await;
result.expect("write + shutdown");
stream.finalize().await.expect("finalize");
let mut staged = stager.locked_staged();
staged.sort_by(|a, b| a.0.cmp(&b.0));
assert_eq!(staged.len(), 2);
assert_eq!(staged[0].1.len(), 4);
assert_eq!(staged[1].1.len(), 2);
}
#[tokio::test(flavor = "multi_thread", worker_threads = 2)]
async fn max_blocks_cap_is_enforced_before_commit() {
let stager = Arc::new(FakeStager::new());
let stream = build_stream_with_max_blocks(stager.clone(), 1, 2, 3);
let payload: Vec<u8> = (0..4).collect();
let (stream, result) = run_write(stream, move |w| w.write_all(&payload)).await;
let write_err = result.expect_err("four single-byte blocks exceed max_blocks=3");
assert_eq!(write_err.to_string(), Error::TooLarge.to_string());
let err = stream
.finalize()
.await
.expect_err("finalize reports the assignment-time cap failure");
assert!(matches!(err, Error::TooLarge), "got: {err:?}");
assert!(
stager.locked_commits().is_empty(),
"cap failure must not commit a partial block list"
);
}
#[tokio::test(flavor = "multi_thread", worker_threads = 2)]
async fn exactly_max_blocks_full_blocks_succeeds() {
let stager = Arc::new(FakeStager::new());
let stream = build_stream_with_max_blocks(stager.clone(), 2, 2, 3);
let payload: Vec<u8> = (0..6).collect();
let (stream, result) = run_write(stream, move |w| w.write_all(&payload)).await;
result.expect("exactly three two-byte blocks are within the cap");
stream
.finalize()
.await
.expect("finalize commits capped write");
let commits = stager.locked_commits();
assert_eq!(commits.len(), 1);
assert_eq!(
commits[0].len(),
3,
"exactly max_blocks block IDs are committed"
);
}
#[tokio::test(flavor = "multi_thread", worker_threads = 2)]
async fn byte_past_max_blocks_fails_on_trailing_partial() {
let stager = Arc::new(FakeStager::new());
let stream = build_stream_with_max_blocks(stager.clone(), 2, 2, 3);
let payload: Vec<u8> = (0..7).collect();
let (stream, result) = run_write(stream, move |w| w.write_all(&payload)).await;
let write_err = result.expect_err("seventh byte requires a fourth block during shutdown");
assert_eq!(write_err.to_string(), Error::TooLarge.to_string());
let err = stream
.finalize()
.await
.expect_err("finalize preserves the shutdown cap failure");
assert!(matches!(err, Error::TooLarge), "got: {err:?}");
assert!(
stager.locked_commits().is_empty(),
"no commit after the boundary failure"
);
}
#[tokio::test(flavor = "multi_thread", worker_threads = 2)]
async fn stage_failure_surfaces_and_skips_commit() {
let stager = Arc::new(FakeStager::failing(1));
let stream = build_stream(stager.clone(), 4, 2);
let payload: Vec<u8> = (0..12).collect();
let (stream, _result) = run_write(stream, move |w| {
drop(w.write_all(&payload));
Ok(())
})
.await;
let err = stream
.finalize()
.await
.expect_err("finalize must report stage failure");
assert!(matches!(err, Error::Io(_)), "got: {err:?}");
let commits = stager.locked_commits();
assert!(commits.is_empty(), "no commit after stage failure");
}
#[tokio::test(flavor = "multi_thread", worker_threads = 2)]
async fn abort_does_not_commit() {
let stager = Arc::new(FakeStager::new());
let stream = build_stream(stager.clone(), 4, 2);
let payload: Vec<u8> = (0..8).collect();
let (stream, result) = run_write(stream, move |w| w.write_all(&payload)).await;
result.expect("write + shutdown");
stream.abort().await.expect("abort succeeds");
let commits = stager.locked_commits();
assert!(commits.is_empty(), "abort skips commit");
}
#[tokio::test(flavor = "multi_thread", worker_threads = 4)]
async fn concurrency_bound_is_respected() {
let stager = Arc::new(FakeStager::new());
let stream = build_stream(stager.clone(), 1, 1);
let payload: Vec<u8> = (0..6).collect();
let (stream, result) = run_write(stream, move |w| w.write_all(&payload)).await;
result.expect("write + shutdown");
stream.finalize().await.expect("finalize");
let observed = stager.max_concurrent_stages.load(Ordering::SeqCst);
assert_eq!(observed, 1, "concurrency=1 caps in-flight at 1");
assert_eq!(stager.stage_call_count.load(Ordering::SeqCst), 6);
}
#[tokio::test(flavor = "multi_thread", worker_threads = 4)]
async fn higher_concurrency_allows_overlap() {
let stager = Arc::new(FakeStager::new());
let stream = build_stream(stager.clone(), 1, 4);
let payload: Vec<u8> = (0..8).collect();
let (stream, result) = run_write(stream, move |w| w.write_all(&payload)).await;
result.expect("write + shutdown");
stream.finalize().await.expect("finalize");
let observed = stager.max_concurrent_stages.load(Ordering::SeqCst);
assert!(observed >= 2, "expected some overlap, got {observed}");
assert!(observed <= 4, "must not exceed configured concurrency");
}
#[test]
fn block_id_round_trip() {
for i in [0_u64, 1, 100, u64::MAX] {
let id = block_id(i);
assert_eq!(id.len(), 8);
let mut buf = [0_u8; 8];
buf.copy_from_slice(&id);
assert_eq!(u64::from_be_bytes(buf), i);
}
}
}