use std::collections::HashSet;
use std::marker::PhantomData;
use std::sync::Arc;
use std::time::{Duration, Instant};
use futures::future::Either;
use tokio::sync::{Semaphore, mpsc, oneshot, watch};
use tokio::task::JoinHandle;
use crate::leader::InstanceLeader;
use crate::object::ObjectBlockOps;
use crate::{BlockId, SequenceHash};
use kvbm_common::LogicalLayoutHandle;
use kvbm_logical::blocks::{BlockMetadata, BlockRegistry, ImmutableBlock};
use kvbm_logical::manager::BlockManager;
use kvbm_physical::transfer::TransferOptions;
use super::batch::{
BatchCollector, BatchConfig, BatchOutputRx, EvalResult, QueuedBlock, TimingTrace, TransferBatch,
};
use super::handle::{TransferId, TransferState, TransferStatus};
use super::pending::PendingTracker;
use super::policy::{EvalContext, OffloadPolicy};
use super::queue::CancellableQueue;
use super::settlement::{
BatchPhaseGuard, PipelineFailure, PipelineFailureKind, PipelineRunGuard,
PipelineSettlementTracker, QueuedBatchGuard,
};
use super::source::{SourceBlock, SourceBlocks};
use crate::object::ObjectLockManager;
#[derive(Clone)]
pub struct PipelineConfig<Src: BlockMetadata, Dst: BlockMetadata> {
pub policies: Vec<Arc<dyn OffloadPolicy<Src>>>,
pub batch_config: BatchConfig,
pub policy_timeout: Duration,
pub auto_chain: bool,
pub eval_input_capacity: usize,
pub batch_input_capacity: usize,
pub transfer_input_capacity: usize,
pub sweep_interval: Duration,
pub skip_transfers: bool,
pub max_concurrent_transfers: usize,
pub ordered_transfer_starts: bool,
pub pending_tracker: Option<Arc<PendingTracker>>,
pub max_concurrent_precondition_awaits: usize,
_marker: PhantomData<(Src, Dst)>,
}
impl<Src: BlockMetadata, Dst: BlockMetadata> Default for PipelineConfig<Src, Dst> {
fn default() -> Self {
Self {
policies: Vec::new(),
batch_config: BatchConfig::default(),
policy_timeout: Duration::from_millis(100),
auto_chain: false,
eval_input_capacity: 128,
batch_input_capacity: 256,
transfer_input_capacity: 8,
sweep_interval: Duration::from_millis(10),
skip_transfers: false,
max_concurrent_transfers: 1,
ordered_transfer_starts: false,
pending_tracker: None,
max_concurrent_precondition_awaits: 8,
_marker: PhantomData,
}
}
}
pub struct PipelineBuilder<Src: BlockMetadata, Dst: BlockMetadata> {
config: PipelineConfig<Src, Dst>,
}
impl<Src: BlockMetadata, Dst: BlockMetadata> PipelineBuilder<Src, Dst> {
pub fn new() -> Self {
Self {
config: PipelineConfig::default(),
}
}
pub fn policy(mut self, policy: Arc<dyn OffloadPolicy<Src>>) -> Self {
self.config.policies.push(policy);
self
}
pub fn batch_size(mut self, size: usize) -> Self {
self.config.batch_config.max_batch_size = size;
self
}
pub fn min_batch_size(mut self, size: usize) -> Self {
self.config.batch_config.min_batch_size = size;
self
}
pub fn flush_interval(mut self, interval: Duration) -> Self {
self.config.batch_config.flush_interval = interval;
self
}
pub fn policy_timeout(mut self, timeout: Duration) -> Self {
self.config.policy_timeout = timeout;
self
}
pub fn auto_chain(mut self, enabled: bool) -> Self {
self.config.auto_chain = enabled;
self
}
pub fn sweep_interval(mut self, interval: Duration) -> Self {
self.config.sweep_interval = interval;
self
}
pub fn skip_transfers(mut self, skip: bool) -> Self {
self.config.skip_transfers = skip;
self
}
pub fn max_concurrent_transfers(mut self, n: usize) -> Self {
self.config.max_concurrent_transfers = n.max(1);
self
}
pub fn ordered_transfer_starts(mut self, ordered: bool) -> Self {
self.config.ordered_transfer_starts = ordered;
self
}
pub fn pending_tracker(mut self, tracker: Arc<PendingTracker>) -> Self {
self.config.pending_tracker = Some(tracker);
self
}
pub fn build(self) -> PipelineConfig<Src, Dst> {
self.config
}
}
impl<Src: BlockMetadata, Dst: BlockMetadata> Default for PipelineBuilder<Src, Dst> {
fn default() -> Self {
Self::new()
}
}
pub(crate) struct PipelineInput<T: BlockMetadata> {
pub(crate) transfer_id: TransferId,
pub(crate) source: SourceBlocks<T>,
pub(crate) state: Arc<std::sync::Mutex<TransferState>>,
}
pub struct PipelineOutput {
pub transfer_id: TransferId,
pub completed_hashes: Vec<SequenceHash>,
}
pub struct ChainOutput<T: BlockMetadata> {
pub transfer_id: TransferId,
pub blocks: Vec<ImmutableBlock<T>>,
#[allow(dead_code)]
pub(crate) state: Arc<std::sync::Mutex<TransferState>>,
}
pub type ChainOutputRx<T> = mpsc::Receiver<ChainOutput<T>>;
pub struct Pipeline<Src: BlockMetadata, Dst: BlockMetadata> {
config: PipelineConfig<Src, Dst>,
pub(crate) eval_queue: Arc<CancellableQueue<PipelineInput<Src>>>,
output_tx: Option<mpsc::Sender<PipelineOutput>>,
chain_rx: Option<ChainOutputRx<Dst>>,
cancel_tx: watch::Sender<HashSet<TransferId>>,
pending_tracker: Arc<PendingTracker>,
pub(crate) settlement: PipelineSettlementTracker,
_task_handles: Vec<JoinHandle<()>>,
_marker: PhantomData<Dst>,
}
impl<Src: BlockMetadata, Dst: BlockMetadata> Pipeline<Src, Dst> {
#[allow(clippy::too_many_arguments)]
pub fn new(
config: PipelineConfig<Src, Dst>,
_registry: Arc<BlockRegistry>,
dst_manager: Arc<BlockManager<Dst>>,
leader: Arc<InstanceLeader>,
src_layout: LogicalLayoutHandle,
dst_layout: LogicalLayoutHandle,
runtime: tokio::runtime::Handle,
) -> Self {
let eval_queue: Arc<CancellableQueue<PipelineInput<Src>>> =
Arc::new(CancellableQueue::new());
let batch_queue: Arc<CancellableQueue<EvalResult<Src>>> = Arc::new(CancellableQueue::new());
let (output_tx, _output_rx) = mpsc::channel(64);
let (cancel_tx, cancel_rx) = watch::channel(HashSet::new());
let (batch_tx, batch_rx) = mpsc::channel(config.transfer_input_capacity);
let (precond_tx, precond_rx) = mpsc::channel(config.transfer_input_capacity);
let (chain_tx, chain_rx) = if config.auto_chain {
let (tx, rx) = mpsc::channel(64);
(Some(tx), Some(rx))
} else {
(None, None)
};
let pending_tracker = config
.pending_tracker
.clone()
.unwrap_or_else(|| Arc::new(PendingTracker::new()));
let settlement = PipelineSettlementTracker::new(config.max_concurrent_transfers);
let evaluator = PolicyEvaluator {
policies: config.policies.clone(),
timeout: config.policy_timeout,
input_queue: eval_queue.clone(),
output_queue: batch_queue.clone(),
cancel_rx: cancel_rx.clone(),
pending_tracker: pending_tracker.clone(),
};
let eval_handle = runtime.spawn(async move {
evaluator.run().await;
});
let collector_input_queue = batch_queue.clone();
let batch_config = config.batch_config.clone();
let collector_cancel_rx = cancel_rx.clone();
let batch_handle = runtime.spawn(async move {
let collector = BatchCollector::new(
batch_config,
collector_input_queue,
batch_tx,
collector_cancel_rx,
);
collector.run().await;
});
let awaiter_leader = leader.clone();
let awaiter_settlement = settlement.clone();
let precond_handle = runtime.spawn(async move {
let awaiter = PreconditionAwaiter {
input_rx: batch_rx,
output_tx: precond_tx,
leader: awaiter_leader,
settlement: awaiter_settlement,
};
awaiter.run().await;
});
let executor = BlockTransferExecutor {
input_rx: precond_rx,
leader,
dst_manager,
src_layout,
dst_layout,
skip_transfers: config.skip_transfers,
max_concurrent_transfers: config.max_concurrent_transfers,
ordered_transfer_starts: config.ordered_transfer_starts,
chain_tx,
settlement: settlement.clone(),
_src_marker: PhantomData::<Src>,
};
let transfer_handle = runtime.spawn(async move {
executor.run().await;
});
let sweeper_queues = vec![eval_queue.clone()];
let sweeper_batch_queue = batch_queue;
let sweeper_interval = config.sweep_interval;
let sweeper_cancel_rx = cancel_rx;
let sweeper_handle = runtime.spawn(async move {
cancel_sweeper(
sweeper_queues,
sweeper_batch_queue,
sweeper_cancel_rx,
sweeper_interval,
)
.await;
});
Self {
config,
eval_queue,
output_tx: Some(output_tx),
chain_rx,
cancel_tx,
pending_tracker,
settlement,
_task_handles: vec![
eval_handle,
batch_handle,
precond_handle,
transfer_handle,
sweeper_handle,
],
_marker: PhantomData,
}
}
pub(crate) fn enqueue(
&self,
transfer_id: TransferId,
source: SourceBlocks<Src>,
state: Arc<std::sync::Mutex<TransferState>>,
) -> bool {
tracing::debug!(%transfer_id, num_blocks = source.len(), "Pipeline: enqueueing blocks");
let input = PipelineInput {
transfer_id,
source,
state,
};
self.eval_queue.push(transfer_id, input)
}
pub fn request_cancel(&self, transfer_id: TransferId) {
self.eval_queue.mark_cancelled(transfer_id);
self.cancel_tx.send_modify(|set| {
set.insert(transfer_id);
});
}
pub fn auto_chain(&self) -> bool {
self.config.auto_chain
}
pub fn output_tx(&self) -> Option<mpsc::Sender<PipelineOutput>> {
self.output_tx.clone()
}
pub fn take_chain_rx(&mut self) -> Option<ChainOutputRx<Dst>> {
self.chain_rx.take()
}
pub fn pending_tracker(&self) -> &Arc<PendingTracker> {
&self.pending_tracker
}
}
#[derive(Clone)]
pub struct ObjectPipelineConfig<Src: BlockMetadata> {
pub policies: Vec<Arc<dyn OffloadPolicy<Src>>>,
pub batch_config: BatchConfig,
pub policy_timeout: Duration,
pub eval_input_capacity: usize,
pub batch_input_capacity: usize,
pub transfer_input_capacity: usize,
pub sweep_interval: Duration,
pub skip_transfers: bool,
pub max_concurrent_transfers: usize,
pub pending_tracker: Option<Arc<PendingTracker>>,
pub max_concurrent_precondition_awaits: usize,
pub lock_manager: Option<Arc<dyn ObjectLockManager>>,
_marker: PhantomData<Src>,
}
impl<Src: BlockMetadata> Default for ObjectPipelineConfig<Src> {
fn default() -> Self {
Self {
policies: Vec::new(),
batch_config: BatchConfig::default(),
policy_timeout: Duration::from_millis(100),
eval_input_capacity: 128,
batch_input_capacity: 256,
transfer_input_capacity: 8,
sweep_interval: Duration::from_millis(10),
skip_transfers: false,
max_concurrent_transfers: 1,
pending_tracker: None,
max_concurrent_precondition_awaits: 8,
lock_manager: None,
_marker: PhantomData,
}
}
}
pub struct ObjectPipelineBuilder<Src: BlockMetadata> {
config: ObjectPipelineConfig<Src>,
}
impl<Src: BlockMetadata> ObjectPipelineBuilder<Src> {
pub fn new() -> Self {
Self {
config: ObjectPipelineConfig::default(),
}
}
pub fn policy(mut self, policy: Arc<dyn OffloadPolicy<Src>>) -> Self {
self.config.policies.push(policy);
self
}
pub fn batch_size(mut self, size: usize) -> Self {
self.config.batch_config.max_batch_size = size;
self
}
pub fn min_batch_size(mut self, size: usize) -> Self {
self.config.batch_config.min_batch_size = size;
self
}
pub fn flush_interval(mut self, interval: Duration) -> Self {
self.config.batch_config.flush_interval = interval;
self
}
pub fn policy_timeout(mut self, timeout: Duration) -> Self {
self.config.policy_timeout = timeout;
self
}
pub fn sweep_interval(mut self, interval: Duration) -> Self {
self.config.sweep_interval = interval;
self
}
pub fn skip_transfers(mut self, skip: bool) -> Self {
self.config.skip_transfers = skip;
self
}
pub fn max_concurrent_transfers(mut self, n: usize) -> Self {
self.config.max_concurrent_transfers = n.max(1);
self
}
pub fn pending_tracker(mut self, tracker: Arc<PendingTracker>) -> Self {
self.config.pending_tracker = Some(tracker);
self
}
pub fn lock_manager(mut self, manager: Arc<dyn ObjectLockManager>) -> Self {
self.config.lock_manager = Some(manager);
self
}
pub fn build(self) -> ObjectPipelineConfig<Src> {
self.config
}
}
impl<Src: BlockMetadata> Default for ObjectPipelineBuilder<Src> {
fn default() -> Self {
Self::new()
}
}
#[allow(dead_code)]
pub struct ObjectPipeline<Src: BlockMetadata> {
config: ObjectPipelineConfig<Src>,
pub(crate) eval_queue: Arc<CancellableQueue<PipelineInput<Src>>>,
output_tx: Option<mpsc::Sender<PipelineOutput>>,
cancel_tx: watch::Sender<HashSet<TransferId>>,
pending_tracker: Arc<PendingTracker>,
pub(crate) settlement: PipelineSettlementTracker,
_task_handles: Vec<JoinHandle<()>>,
}
impl<Src: BlockMetadata> ObjectPipeline<Src> {
#[allow(clippy::too_many_arguments)]
pub fn new(
config: ObjectPipelineConfig<Src>,
object_ops: Arc<dyn ObjectBlockOps>,
src_layout: LogicalLayoutHandle,
leader: Arc<InstanceLeader>,
runtime: tokio::runtime::Handle,
) -> Self {
let eval_queue: Arc<CancellableQueue<PipelineInput<Src>>> =
Arc::new(CancellableQueue::new());
let batch_queue: Arc<CancellableQueue<EvalResult<Src>>> = Arc::new(CancellableQueue::new());
let (output_tx, _output_rx) = mpsc::channel(64);
let (cancel_tx, cancel_rx) = watch::channel(HashSet::new());
let (batch_tx, batch_rx) = mpsc::channel(config.transfer_input_capacity);
let (precond_tx, precond_rx) = mpsc::channel(config.transfer_input_capacity);
let pending_tracker = config
.pending_tracker
.clone()
.unwrap_or_else(|| Arc::new(PendingTracker::new()));
let settlement = PipelineSettlementTracker::new(config.max_concurrent_transfers);
let evaluator = PolicyEvaluator {
policies: config.policies.clone(),
timeout: config.policy_timeout,
input_queue: eval_queue.clone(),
output_queue: batch_queue.clone(),
cancel_rx: cancel_rx.clone(),
pending_tracker: pending_tracker.clone(),
};
let eval_handle = runtime.spawn(async move {
evaluator.run().await;
});
let collector_input_queue = batch_queue.clone();
let batch_config = config.batch_config.clone();
let collector_cancel_rx = cancel_rx.clone();
let batch_handle = runtime.spawn(async move {
let collector = BatchCollector::new(
batch_config,
collector_input_queue,
batch_tx,
collector_cancel_rx,
);
collector.run().await;
});
let awaiter_leader = leader.clone();
let awaiter_settlement = settlement.clone();
let precond_handle = runtime.spawn(async move {
let awaiter = PreconditionAwaiter {
input_rx: batch_rx,
output_tx: precond_tx,
leader: awaiter_leader,
settlement: awaiter_settlement,
};
awaiter.run().await;
});
let executor = ObjectTransferExecutor::new(
precond_rx,
object_ops,
src_layout,
config.skip_transfers,
config.max_concurrent_transfers,
config.lock_manager.clone(),
settlement.clone(),
);
let transfer_handle = runtime.spawn(async move {
executor.run().await;
});
let sweeper_queues = vec![eval_queue.clone()];
let sweeper_batch_queue = batch_queue;
let sweeper_interval = config.sweep_interval;
let sweeper_cancel_rx = cancel_rx;
let sweeper_handle = runtime.spawn(async move {
cancel_sweeper(
sweeper_queues,
sweeper_batch_queue,
sweeper_cancel_rx,
sweeper_interval,
)
.await;
});
Self {
config,
eval_queue,
output_tx: Some(output_tx),
cancel_tx,
pending_tracker,
settlement,
_task_handles: vec![
eval_handle,
batch_handle,
precond_handle,
transfer_handle,
sweeper_handle,
],
}
}
pub(crate) fn enqueue(
&self,
transfer_id: TransferId,
source: SourceBlocks<Src>,
state: Arc<std::sync::Mutex<TransferState>>,
) -> bool {
tracing::debug!(%transfer_id, num_blocks = source.len(), "ObjectPipeline: enqueueing blocks");
let input = PipelineInput {
transfer_id,
source,
state,
};
self.eval_queue.push(transfer_id, input)
}
pub fn request_cancel(&self, transfer_id: TransferId) {
self.eval_queue.mark_cancelled(transfer_id);
self.cancel_tx.send_modify(|set| {
set.insert(transfer_id);
});
}
#[allow(dead_code)]
pub fn output_tx(&self) -> Option<mpsc::Sender<PipelineOutput>> {
self.output_tx.clone()
}
pub fn pending_tracker(&self) -> &Arc<PendingTracker> {
&self.pending_tracker
}
}
async fn cancel_sweeper<Src: BlockMetadata>(
input_queues: Vec<Arc<CancellableQueue<PipelineInput<Src>>>>,
batch_queue: Arc<CancellableQueue<EvalResult<Src>>>,
mut cancel_rx: watch::Receiver<HashSet<TransferId>>,
interval: Duration,
) {
let mut ticker = tokio::time::interval(interval);
ticker.set_missed_tick_behavior(tokio::time::MissedTickBehavior::Skip);
loop {
tokio::select! {
_ = ticker.tick() => {
for queue in &input_queues {
let removed = queue.sweep();
if removed > 0 {
tracing::debug!("Sweeper removed {} cancelled input items", removed);
}
}
let batch_removed = batch_queue.sweep();
if batch_removed > 0 {
tracing::debug!("Sweeper removed {} cancelled batch items", batch_removed);
}
}
result = cancel_rx.changed() => {
if result.is_err() {
break;
}
for queue in &input_queues {
queue.sweep();
}
batch_queue.sweep();
}
}
}
}
struct PolicyEvaluator<T: BlockMetadata> {
policies: Vec<Arc<dyn OffloadPolicy<T>>>,
timeout: Duration,
input_queue: Arc<CancellableQueue<PipelineInput<T>>>,
output_queue: Arc<CancellableQueue<EvalResult<T>>>,
cancel_rx: watch::Receiver<HashSet<TransferId>>,
pending_tracker: Arc<PendingTracker>,
}
impl<T: BlockMetadata> PolicyEvaluator<T> {
async fn run(mut self) {
loop {
while let Some(item) = self.input_queue.pop_valid() {
self.evaluate(item.data).await;
}
tokio::select! {
_ = self.input_queue.notified() => {}
result = self.cancel_rx.changed() => {
if result.is_err() {
break;
}
}
}
}
}
async fn evaluate(&self, input: PipelineInput<T>) {
nvtx_range!("offload::policy");
let transfer_id = input.transfer_id;
let total_blocks = input.source.len();
{
let mut state = input.state.lock().unwrap();
state.total_expected_blocks = total_blocks;
}
{
let state = input.state.lock().unwrap();
if state.is_cancel_requested() {
drop(state); tracing::debug!(%transfer_id, "Transfer cancelled before evaluation");
let mut state = input.state.lock().unwrap();
state.set_cancelled();
return;
}
}
let mut passed = Vec::new();
let mut filtered = Vec::new();
match input.source {
SourceBlocks::External(external_blocks) => {
for ext in external_blocks {
if self.check_cancelled(&input.state, transfer_id) {
return;
}
let ctx = EvalContext::from_external(ext.block_id, ext.sequence_hash);
let pass = self.evaluate_policies(&ctx).await;
if pass {
let pending_guard = self.pending_tracker.guard(ext.sequence_hash);
passed.push(QueuedBlock {
transfer_id,
block_id: Some(ext.block_id),
sequence_hash: ext.sequence_hash,
source: SourceBlock::External(ext),
state: input.state.clone(),
pending_guard: Some(pending_guard),
});
} else {
filtered.push(ext.block_id);
}
}
tracing::debug!(%transfer_id, passed = passed.len(), filtered = filtered.len(), "External blocks evaluated");
}
SourceBlocks::Strong(strong_blocks) => {
for block in strong_blocks {
if self.check_cancelled(&input.state, transfer_id) {
return;
}
let ctx = EvalContext::new(block);
let pass = self.evaluate_policies(&ctx).await;
if pass {
let block = ctx.block.expect("Strong block context always has block");
let pending_guard = self.pending_tracker.guard(ctx.sequence_hash);
passed.push(QueuedBlock {
transfer_id,
block_id: Some(ctx.block_id),
sequence_hash: ctx.sequence_hash,
source: SourceBlock::Strong(block),
state: input.state.clone(),
pending_guard: Some(pending_guard),
});
} else {
filtered.push(ctx.block_id);
}
}
}
SourceBlocks::Weak(weak_blocks) => {
for weak in weak_blocks {
if self.check_cancelled(&input.state, transfer_id) {
return;
}
let sequence_hash = weak.sequence_hash();
let ctx = EvalContext::from_weak(BlockId::default(), sequence_hash);
let pass = self.evaluate_policies(&ctx).await;
if pass {
let pending_guard = self.pending_tracker.guard(sequence_hash);
passed.push(QueuedBlock {
transfer_id,
block_id: None, sequence_hash,
source: SourceBlock::Weak(weak),
state: input.state.clone(),
pending_guard: Some(pending_guard),
});
} else {
tracing::debug!(%transfer_id, ?sequence_hash, "Weak block filtered by policy");
}
}
}
}
{
let state = input.state.lock().unwrap();
if state.is_cancel_requested() {
drop(state);
tracing::debug!(%transfer_id, "Transfer cancelled after evaluation");
let mut state = input.state.lock().unwrap();
state.set_cancelled();
return;
}
}
tracing::debug!(%transfer_id, passed = passed.len(), filtered = filtered.len(), "Policy evaluation complete");
{
let mut state = input.state.lock().unwrap();
state.add_passed(passed.iter().filter_map(|b| b.block_id));
state.add_filtered(filtered.iter().copied());
state.set_status(TransferStatus::Queued);
}
if passed.is_empty() {
tracing::debug!(%transfer_id, "All blocks filtered, completing transfer");
let mut state = input.state.lock().unwrap();
state.set_complete();
return;
}
let result = EvalResult {
transfer_id,
passed_blocks: passed,
filtered_ids: filtered,
state: input.state,
};
if !self.output_queue.push(transfer_id, result) {
tracing::debug!(%transfer_id, "Push to output queue failed (cancelled)");
}
}
fn check_cancelled(
&self,
state: &Arc<std::sync::Mutex<TransferState>>,
transfer_id: TransferId,
) -> bool {
let state_guard = state.lock().unwrap();
if state_guard.is_cancel_requested() {
drop(state_guard);
tracing::debug!(%transfer_id, "Transfer cancelled mid-evaluation");
let mut state_guard = state.lock().unwrap();
state_guard.set_cancelled();
true
} else {
false
}
}
async fn evaluate_policies(&self, ctx: &EvalContext<T>) -> bool {
for policy in &self.policies {
let eval_future = policy.evaluate(ctx);
let timed_result = tokio::time::timeout(self.timeout, async {
match eval_future {
Either::Left(ready) => ready.await,
Either::Right(boxed) => boxed.await,
}
})
.await;
match timed_result {
Ok(Ok(true)) => continue,
Ok(Ok(false)) => return false,
Ok(Err(e)) => {
tracing::warn!("Policy {} error: {}", policy.name(), e);
return false;
}
Err(_) => {
tracing::warn!("Policy {} timed out", policy.name());
return false;
}
}
}
true
}
}
pub struct ResolvedBlock<T: BlockMetadata> {
pub transfer_id: TransferId,
pub block_id: BlockId,
pub sequence_hash: SequenceHash,
#[allow(dead_code)]
pub guard: Option<ImmutableBlock<T>>,
pub(crate) state: Arc<std::sync::Mutex<TransferState>>,
}
pub struct ResolvedBatch<T: BlockMetadata> {
pub blocks: Vec<ResolvedBlock<T>>,
#[allow(dead_code)]
pub evicted: Vec<SequenceHash>,
pub timing: TimingTrace,
}
impl<T: BlockMetadata> ResolvedBatch<T> {
pub fn is_empty(&self) -> bool {
self.blocks.is_empty()
}
#[allow(dead_code)]
pub fn len(&self) -> usize {
self.blocks.len()
}
}
pub fn upgrade_batch<T: BlockMetadata>(batch: TransferBatch<T>) -> ResolvedBatch<T> {
let mut resolved: Vec<ResolvedBlock<T>> = Vec::with_capacity(batch.len());
let mut evicted: Vec<SequenceHash> = Vec::new();
let mut timing = batch.timing;
timing.mark_transfer_start();
for queued in batch.blocks {
match queued.source {
SourceBlock::Strong(block) => {
resolved.push(ResolvedBlock {
transfer_id: queued.transfer_id,
block_id: block.block_id(),
sequence_hash: queued.sequence_hash,
guard: Some(block),
state: queued.state,
});
}
SourceBlock::External(ext) => {
resolved.push(ResolvedBlock {
transfer_id: queued.transfer_id,
block_id: ext.block_id,
sequence_hash: ext.sequence_hash,
guard: None,
state: queued.state,
});
}
SourceBlock::Weak(weak) => match weak.upgrade() {
Some(block) => {
resolved.push(ResolvedBlock {
transfer_id: queued.transfer_id,
block_id: block.block_id(),
sequence_hash: queued.sequence_hash,
guard: Some(block),
state: queued.state,
});
}
None => {
tracing::debug!(
sequence_hash = ?queued.sequence_hash,
"Weak block evicted before transfer"
);
evicted.push(queued.sequence_hash);
}
},
}
}
ResolvedBatch {
blocks: resolved,
evicted,
timing,
}
}
struct PreconditionAwaiter<T: BlockMetadata> {
input_rx: BatchOutputRx<T>,
output_tx: mpsc::Sender<TransferBatch<T>>,
leader: Arc<InstanceLeader>,
settlement: PipelineSettlementTracker,
}
impl<T: BlockMetadata> PreconditionAwaiter<T> {
async fn run(mut self) {
while let Some(mut batch) = self.input_rx.recv().await {
let output_tx = self.output_tx.clone();
let nova = self.leader.messenger().clone();
let settlement = self.settlement.clone();
tokio::spawn(async move {
nvtx_range!("offload::precondition");
if let Some(event_handle) = batch.precondition {
tracing::debug!(?event_handle, "Awaiting precondition for batch");
let awaiter_result = nova.events().awaiter(event_handle);
match awaiter_result {
Ok(awaiter) => {
match tokio::time::timeout(Duration::from_secs(300), awaiter).await {
Ok(Ok(())) => {
tracing::debug!(?event_handle, "Precondition satisfied");
}
Ok(Err(poison)) => {
tracing::error!(
?event_handle,
?poison,
"Precondition poisoned, marking all blocks as failed"
);
for queued in batch.blocks {
let mut state = queued.state.lock().unwrap();
state.set_error(format!(
"precondition poisoned: {:?}",
poison
));
}
return;
}
Err(_) => {
tracing::error!(
?event_handle,
"Precondition timeout after 30s"
);
for queued in batch.blocks {
let mut state = queued.state.lock().unwrap();
state.set_error("precondition timeout".to_string());
}
return;
}
}
}
Err(e) => {
tracing::error!(?event_handle, ?e, "Failed to create awaiter");
for queued in batch.blocks {
let mut state = queued.state.lock().unwrap();
state.set_error(format!("failed to create awaiter: {}", e));
}
return;
}
}
}
batch.timing.mark_precondition_complete();
let queued = QueuedBatchGuard::new(settlement);
if let Err(e) = output_tx.send(batch).await {
tracing::error!("Failed to forward batch after precondition: {}", e);
queued.finish_failure(PipelineFailure::new(
PipelineFailureKind::Shutdown,
"transfer executor input channel closed",
));
} else {
queued.sent();
}
});
}
}
}
struct BlockTransferExecutor<Src: BlockMetadata, Dst: BlockMetadata> {
input_rx: BatchOutputRx<Src>,
leader: Arc<InstanceLeader>,
dst_manager: Arc<BlockManager<Dst>>,
src_layout: LogicalLayoutHandle,
dst_layout: LogicalLayoutHandle,
skip_transfers: bool,
max_concurrent_transfers: usize,
ordered_transfer_starts: bool,
chain_tx: Option<mpsc::Sender<ChainOutput<Dst>>>,
settlement: PipelineSettlementTracker,
_src_marker: PhantomData<Src>,
}
struct SharedBlockExecutorState<Dst: BlockMetadata> {
leader: Arc<InstanceLeader>,
dst_manager: Arc<BlockManager<Dst>>,
src_layout: LogicalLayoutHandle,
dst_layout: LogicalLayoutHandle,
skip_transfers: bool,
chain_tx: Option<mpsc::Sender<ChainOutput<Dst>>>,
}
struct OrderedTransferStart {
predecessor: Option<oneshot::Receiver<()>>,
successor: Option<oneshot::Sender<()>>,
}
impl OrderedTransferStart {
fn next(predecessor: &mut Option<oneshot::Receiver<()>>) -> Self {
let (successor, next_predecessor) = oneshot::channel();
Self {
predecessor: predecessor.replace(next_predecessor),
successor: Some(successor),
}
}
async fn wait(&mut self) {
if let Some(predecessor) = self.predecessor.take() {
let _ = predecessor.await;
}
}
fn release(&mut self) {
if let Some(successor) = self.successor.take() {
let _ = successor.send(());
}
}
}
impl Drop for OrderedTransferStart {
fn drop(&mut self) {
self.release();
}
}
impl<Src: BlockMetadata, Dst: BlockMetadata> BlockTransferExecutor<Src, Dst> {
async fn run(mut self) {
let transfer_semaphore = Arc::new(Semaphore::new(self.max_concurrent_transfers));
let prepare_semaphore = Arc::new(Semaphore::new(1));
let shared = Arc::new(SharedBlockExecutorState {
leader: self.leader.clone(),
dst_manager: self.dst_manager.clone(),
src_layout: self.src_layout,
dst_layout: self.dst_layout,
skip_transfers: self.skip_transfers,
chain_tx: self.chain_tx.take(),
});
let settlement = self.settlement.clone();
let run_guard = PipelineRunGuard::new(settlement.clone(), "block transfer executor");
let mut transfer_start_predecessor = None;
while let Some(batch) = self.input_rx.recv().await {
if batch.is_empty() {
settlement.discard_queued();
continue;
}
let prepare_permit = prepare_semaphore.clone().acquire_owned().await;
if prepare_permit.is_err() {
settlement.fail_queued(PipelineFailure::new(
PipelineFailureKind::Shutdown,
"block preparation semaphore closed",
));
break; }
let prepare_permit = prepare_permit.unwrap();
let upgraded = upgrade_batch(batch);
drop(prepare_permit);
if upgraded.is_empty() {
tracing::debug!("All blocks in batch evicted, skipping transfer");
settlement.discard_queued();
continue;
}
let transfer_permit = transfer_semaphore.clone().acquire_owned().await;
if transfer_permit.is_err() {
settlement.fail_queued(PipelineFailure::new(
PipelineFailureKind::Shutdown,
"block transfer semaphore closed",
));
break; }
let transfer_permit = transfer_permit.unwrap();
let shared_clone = shared.clone();
let mut phase = BatchPhaseGuard::starting(settlement.clone(), transfer_permit);
let mut transfer_start = self
.ordered_transfer_starts
.then(|| OrderedTransferStart::next(&mut transfer_start_predecessor));
tokio::spawn(async move {
if let Some(transfer_start) = &mut transfer_start {
transfer_start.wait().await;
}
match Self::execute_transfer(
&shared_clone,
upgraded,
&mut phase,
transfer_start.as_mut(),
)
.await
{
Ok(()) => phase.finish_success(),
Err(error) => {
tracing::error!("BlockTransferExecutor: transfer failed: {}", error);
phase.finish_failure(PipelineFailure::new(
PipelineFailureKind::Executor,
error.to_string(),
));
}
}
});
}
let _ = transfer_semaphore
.acquire_many(self.max_concurrent_transfers as u32)
.await;
run_guard.finish_shutdown();
}
fn fail_transfer_states(
transfer_states: &std::collections::HashMap<
TransferId,
(Arc<std::sync::Mutex<TransferState>>, Vec<BlockId>),
>,
error: &anyhow::Error,
) {
let message = format!("block transfer failed: {error}");
for (state, block_ids) in transfer_states.values() {
let mut state = state.lock().unwrap();
state.mark_failed(block_ids.iter().copied());
state.set_error(message.clone());
}
}
async fn execute_transfer(
shared: &SharedBlockExecutorState<Dst>,
mut batch: ResolvedBatch<Src>,
phase: &mut BatchPhaseGuard,
mut transfer_start: Option<&mut OrderedTransferStart>,
) -> anyhow::Result<()> {
nvtx_range!("offload::transfer");
if batch.is_empty() {
return Ok(());
}
let resolved = &batch.blocks;
let src_block_ids: Vec<BlockId> = resolved.iter().map(|b| b.block_id).collect();
let sequence_hashes: Vec<SequenceHash> = resolved.iter().map(|b| b.sequence_hash).collect();
let mut transfer_states: std::collections::HashMap<
TransferId,
(Arc<std::sync::Mutex<TransferState>>, Vec<BlockId>),
> = std::collections::HashMap::new();
for block in resolved {
transfer_states
.entry(block.transfer_id)
.or_insert_with(|| (block.state.clone(), Vec::new()))
.1
.push(block.block_id);
}
if !shared.skip_transfers {
let Some(dst_blocks) = shared.dst_manager.allocate_blocks(resolved.len()) else {
let error =
anyhow::anyhow!("failed to allocate {} destination blocks", resolved.len());
Self::fail_transfer_states(&transfer_states, &error);
return Err(error);
};
let dst_block_ids: Vec<BlockId> = dst_blocks.iter().map(|b| b.block_id()).collect();
let start_xfer = Instant::now();
let notification = match shared.leader.execute_local_transfer(
shared.src_layout,
shared.dst_layout,
src_block_ids.clone(),
dst_block_ids.clone(),
TransferOptions::default(),
) {
Ok(notification) => notification,
Err(error) => {
Self::fail_transfer_states(&transfer_states, &error);
return Err(error);
}
};
for (state, block_ids) in transfer_states.values() {
let mut state_guard = state.lock().unwrap();
state_guard.set_status(TransferStatus::Transferring);
state_guard.mark_in_flight(block_ids.iter().copied());
}
phase.mark_in_flight();
if let Some(transfer_start) = &mut transfer_start {
transfer_start.release();
}
if let Err(error) = notification.await {
Self::fail_transfer_states(&transfer_states, &error);
return Err(error);
}
phase.mark_settling();
let end_xfer = Instant::now();
let completed_blocks = dst_blocks
.into_iter()
.zip(sequence_hashes.iter())
.map(|(dst_block, seq_hash)| {
dst_block
.stage(*seq_hash, shared.dst_manager.block_size())
.expect("block size mismatch")
})
.collect();
let registered_blocks: Vec<ImmutableBlock<Dst>> =
shared.dst_manager.register_blocks(completed_blocks);
let registration_timepoint = Instant::now();
let unique_transfer_ids: std::collections::HashSet<_> =
resolved.iter().map(|b| b.transfer_id).collect();
let policy_ms = batch
.timing
.policy_duration()
.map(|d| d.as_millis() as u64)
.unwrap_or(0);
let precondition_ms = batch
.timing
.precondition_duration()
.map(|d| d.as_millis() as u64)
.unwrap_or(0);
let total_ms = batch
.timing
.total_duration()
.map(|d| d.as_millis() as u64)
.unwrap_or(0);
tracing::info!(
blocks = resolved.len(),
containers = unique_transfer_ids.len(),
policy_ms,
precondition_ms,
xfer_ms = end_xfer.duration_since(start_xfer).as_millis() as u64,
registration_ms =
registration_timepoint.duration_since(end_xfer).as_millis() as u64,
total_ms,
src = std::any::type_name::<Src>(),
dst = std::any::type_name::<Dst>(),
"Batch transfer complete"
);
if let Some(chain_tx) = &shared.chain_tx {
#[allow(clippy::type_complexity)]
let mut chain_outputs: std::collections::HashMap<
TransferId,
(
Arc<std::sync::Mutex<TransferState>>,
Vec<ImmutableBlock<Dst>>,
),
> = std::collections::HashMap::new();
for (registered, resolved_block) in
registered_blocks.into_iter().zip(resolved.iter())
{
chain_outputs
.entry(resolved_block.transfer_id)
.or_insert_with(|| (resolved_block.state.clone(), Vec::new()))
.1
.push(registered);
}
for (transfer_id, (state, blocks)) in chain_outputs {
let output = ChainOutput {
transfer_id,
blocks,
state,
};
if chain_tx.send(output).await.is_err() {
tracing::warn!(
%transfer_id,
"Chain channel closed, downstream pipeline unavailable"
);
} else {
tracing::debug!(
%transfer_id,
"Sent blocks to chain output for downstream processing"
);
}
}
}
} else {
phase.mark_in_flight();
if let Some(transfer_start) = &mut transfer_start {
transfer_start.release();
}
phase.mark_settling();
for (state, block_ids) in transfer_states.values() {
let mut state_guard = state.lock().unwrap();
state_guard.set_status(TransferStatus::Transferring);
state_guard.mark_in_flight(block_ids.iter().copied());
}
}
batch.timing.mark_transfer_complete();
for (transfer_id, (state, block_ids)) in transfer_states {
let mut state_guard = state.lock().unwrap();
state_guard.mark_completed(block_ids);
let progress = state_guard.progress_counts();
let total = progress.passed + state_guard.filtered_out.len();
let done = progress.completed + state_guard.filtered_out.len();
tracing::debug!(
%transfer_id,
total,
done,
passed = progress.passed,
filtered = state_guard.filtered_out.len(),
completed = progress.completed,
"Transfer batch progress"
);
if done >= total && total > 0 {
state_guard.set_complete();
}
}
Ok(())
}
}
pub struct ObjectTransferExecutor<Src: BlockMetadata> {
input_rx: BatchOutputRx<Src>,
object_ops: Arc<dyn ObjectBlockOps>,
src_layout: LogicalLayoutHandle,
skip_transfers: bool,
max_concurrent_transfers: usize,
lock_manager: Option<Arc<dyn ObjectLockManager>>,
settlement: PipelineSettlementTracker,
}
struct SharedObjectExecutorState {
object_ops: Arc<dyn ObjectBlockOps>,
src_layout: LogicalLayoutHandle,
skip_transfers: bool,
lock_manager: Option<Arc<dyn ObjectLockManager>>,
}
impl<Src: BlockMetadata> ObjectTransferExecutor<Src> {
#[allow(dead_code)]
pub fn new(
input_rx: BatchOutputRx<Src>,
object_ops: Arc<dyn ObjectBlockOps>,
src_layout: LogicalLayoutHandle,
skip_transfers: bool,
max_concurrent_transfers: usize,
lock_manager: Option<Arc<dyn ObjectLockManager>>,
settlement: PipelineSettlementTracker,
) -> Self {
Self {
input_rx,
object_ops,
src_layout,
skip_transfers,
max_concurrent_transfers,
lock_manager,
settlement,
}
}
pub async fn run(mut self) {
let transfer_semaphore = Arc::new(Semaphore::new(self.max_concurrent_transfers));
let prepare_semaphore = Arc::new(Semaphore::new(1));
let shared = Arc::new(SharedObjectExecutorState {
object_ops: self.object_ops.clone(),
src_layout: self.src_layout,
skip_transfers: self.skip_transfers,
lock_manager: self.lock_manager.clone(),
});
let settlement = self.settlement.clone();
let run_guard = PipelineRunGuard::new(settlement.clone(), "object transfer executor");
while let Some(batch) = self.input_rx.recv().await {
if batch.is_empty() {
settlement.discard_queued();
continue;
}
let prepare_permit = prepare_semaphore.clone().acquire_owned().await;
if prepare_permit.is_err() {
settlement.fail_queued(PipelineFailure::new(
PipelineFailureKind::Shutdown,
"object preparation semaphore closed",
));
break; }
let prepare_permit = prepare_permit.unwrap();
let upgraded = upgrade_batch(batch);
drop(prepare_permit);
if upgraded.is_empty() {
tracing::debug!("All blocks in batch evicted, skipping object transfer");
settlement.discard_queued();
continue;
}
let transfer_permit = transfer_semaphore.clone().acquire_owned().await;
if transfer_permit.is_err() {
settlement.fail_queued(PipelineFailure::new(
PipelineFailureKind::Shutdown,
"object transfer semaphore closed",
));
break; }
let transfer_permit = transfer_permit.unwrap();
let shared_clone = shared.clone();
let mut phase = BatchPhaseGuard::starting(settlement.clone(), transfer_permit);
tokio::spawn(async move {
match Self::execute_transfer(&shared_clone, upgraded, &mut phase).await {
Ok(()) => phase.finish_success(),
Err(error) => {
tracing::error!("ObjectTransferExecutor: transfer failed: {}", error);
phase.finish_failure(PipelineFailure::new(
PipelineFailureKind::Executor,
error.to_string(),
));
}
}
});
}
let _ = transfer_semaphore
.acquire_many(self.max_concurrent_transfers as u32)
.await;
run_guard.finish_shutdown();
}
async fn execute_transfer(
shared: &SharedObjectExecutorState,
mut batch: ResolvedBatch<Src>,
phase: &mut BatchPhaseGuard,
) -> anyhow::Result<()> {
nvtx_range!("offload::transfer");
if batch.is_empty() {
return Ok(());
}
let resolved = &batch.blocks;
let keys: Vec<SequenceHash> = resolved.iter().map(|b| b.sequence_hash).collect();
let block_ids: Vec<BlockId> = resolved.iter().map(|b| b.block_id).collect();
let mut transfer_states: std::collections::HashMap<
TransferId,
(Arc<std::sync::Mutex<TransferState>>, Vec<BlockId>),
> = std::collections::HashMap::new();
for block in resolved {
transfer_states
.entry(block.transfer_id)
.or_insert_with(|| (block.state.clone(), Vec::new()))
.1
.push(block.block_id);
}
let mut successful_hashes: Vec<SequenceHash> = Vec::new();
if !shared.skip_transfers {
let mut put = shared
.object_ops
.put_blocks(keys.clone(), shared.src_layout, block_ids);
let mut first_poll = true;
let results = std::future::poll_fn(|cx| {
if first_poll {
first_poll = false;
phase.mark_in_flight();
}
put.as_mut().poll(cx)
})
.await;
phase.mark_settling();
if results.len() != keys.len() {
tracing::error!(
expected = keys.len(),
actual = results.len(),
"put_blocks returned mismatched result count"
);
for (_transfer_id, (state, block_ids)) in transfer_states {
let mut state_guard = state.lock().unwrap();
state_guard.mark_failed(block_ids);
state_guard
.set_error("put_blocks returned mismatched result count".to_string());
}
return Ok(());
}
let mut success_count = 0;
let mut fail_count = 0;
for result in results {
match result {
Ok(hash) => {
success_count += 1;
successful_hashes.push(hash);
}
Err(hash) => {
fail_count += 1;
tracing::warn!(?hash, "Failed to transfer block to object storage");
}
}
}
if fail_count > 0 {
tracing::warn!(
success = success_count,
failed = fail_count,
"Object transfer partially failed"
);
} else {
tracing::debug!(
num_blocks = success_count,
"Successfully transferred blocks to object storage"
);
}
if let Some(lock_manager) = &shared.lock_manager {
for hash in &successful_hashes {
if let Err(e) = lock_manager.create_meta(*hash).await {
tracing::error!(?hash, error = %e, "Failed to create meta file");
}
if let Err(e) = lock_manager.release_lock(*hash).await {
tracing::error!(?hash, error = %e, "Failed to release lock");
}
}
tracing::debug!(
num_blocks = successful_hashes.len(),
"Created meta files and released locks"
);
}
} else {
phase.mark_in_flight();
phase.mark_settling();
if let Some(lock_manager) = &shared.lock_manager {
for hash in &keys {
if let Err(e) = lock_manager.create_meta(*hash).await {
tracing::error!(?hash, error = %e, "Failed to create meta file");
}
if let Err(e) = lock_manager.release_lock(*hash).await {
tracing::error!(?hash, error = %e, "Failed to release lock");
}
}
}
}
batch.timing.mark_transfer_complete();
let unique_transfer_ids: std::collections::HashSet<_> =
resolved.iter().map(|b| b.transfer_id).collect();
let policy_ms = batch
.timing
.policy_duration()
.map(|d| d.as_millis() as u64)
.unwrap_or(0);
let precondition_ms = batch
.timing
.precondition_duration()
.map(|d| d.as_millis() as u64)
.unwrap_or(0);
let transfer_ms = batch
.timing
.transfer_duration()
.map(|d| d.as_millis() as u64)
.unwrap_or(0);
let total_ms = batch
.timing
.total_duration()
.map(|d| d.as_millis() as u64)
.unwrap_or(0);
tracing::info!(
blocks = resolved.len(),
containers = unique_transfer_ids.len(),
policy_ms,
precondition_ms,
transfer_ms,
total_ms,
src = std::any::type_name::<Src>(),
dst = "G4-object",
"Object batch transfer complete"
);
let block_to_hash: std::collections::HashMap<BlockId, SequenceHash> = resolved
.iter()
.map(|b| (b.block_id, b.sequence_hash))
.collect();
let success_set: std::collections::HashSet<SequenceHash> =
successful_hashes.into_iter().collect();
debug_assert_eq!(
block_to_hash.len(),
resolved.len(),
"duplicate BlockId in batch — block_to_hash would lose entries"
);
debug_assert_eq!(
resolved
.iter()
.map(|b| b.sequence_hash)
.collect::<std::collections::HashSet<_>>()
.len(),
resolved.len(),
"duplicate SequenceHash in batch — hash-based success correlation is ambiguous"
);
for (transfer_id, (state, block_ids)) in transfer_states {
let mut state_guard = state.lock().unwrap();
if shared.skip_transfers {
state_guard.mark_completed(block_ids);
} else {
let (succeeded, failed): (Vec<_>, Vec<_>) = block_ids.into_iter().partition(|id| {
block_to_hash
.get(id)
.is_some_and(|h| success_set.contains(h))
});
state_guard.mark_completed(succeeded);
if !failed.is_empty() {
tracing::warn!(
%transfer_id,
failed_count = failed.len(),
"Marking blocks as failed in transfer state"
);
state_guard.mark_failed(failed);
}
}
let progress = state_guard.progress_counts();
let total = progress.passed + state_guard.filtered_out.len();
let done = progress.settled() + state_guard.filtered_out.len();
tracing::debug!(
%transfer_id,
total,
done,
passed = progress.passed,
filtered = state_guard.filtered_out.len(),
completed = progress.completed,
failed = progress.failed,
"Object transfer batch progress"
);
if done >= total && total > 0 {
let failed_count = progress.failed;
if failed_count == 0 {
state_guard.set_complete();
} else {
state_guard.set_error(format!(
"{failed_count} blocks failed to transfer to object storage",
));
}
}
}
Ok(())
}
}
#[cfg(test)]
mod tests {
use futures::FutureExt;
use super::*;
#[test]
fn test_pipeline_builder() {
let config = PipelineBuilder::<(), ()>::new()
.batch_size(32)
.min_batch_size(8)
.policy_timeout(Duration::from_millis(50))
.auto_chain(true)
.sweep_interval(Duration::from_millis(5))
.build();
assert_eq!(config.batch_config.max_batch_size, 32);
assert_eq!(config.batch_config.min_batch_size, 8);
assert_eq!(config.policy_timeout, Duration::from_millis(50));
assert!(config.auto_chain);
assert_eq!(config.sweep_interval, Duration::from_millis(5));
}
#[test]
fn test_pipeline_config_default() {
let config = PipelineConfig::<(), ()>::default();
assert!(config.policies.is_empty());
assert!(!config.auto_chain);
assert_eq!(config.sweep_interval, Duration::from_millis(10));
assert!(!config.ordered_transfer_starts);
}
#[tokio::test]
async fn ordered_transfer_starts_ignore_task_schedule_order() {
let observed = Arc::new(tokio::sync::Mutex::new(Vec::new()));
let mut predecessor = None;
let starts = (0..3)
.map(|_| OrderedTransferStart::next(&mut predecessor))
.collect::<Vec<_>>();
let mut tasks = Vec::new();
for (index, mut start) in starts.into_iter().enumerate().rev() {
let observed = observed.clone();
tasks.push(tokio::spawn(async move {
start.wait().await;
observed.lock().await.push(index);
start.release();
}));
}
for task in tasks {
task.await.unwrap();
}
assert_eq!(*observed.lock().await, vec![0, 1, 2]);
}
struct FailableObjectBlockOps {
fail_hashes: std::collections::HashSet<SequenceHash>,
}
impl crate::object::ObjectBlockOps for FailableObjectBlockOps {
fn has_blocks(
&self,
keys: Vec<SequenceHash>,
) -> futures::future::BoxFuture<'static, Vec<(SequenceHash, Option<usize>)>> {
Box::pin(async move { keys.into_iter().map(|h| (h, Some(1))).collect() })
}
fn put_blocks(
&self,
keys: Vec<SequenceHash>,
_layout: LogicalLayoutHandle,
_block_ids: Vec<BlockId>,
) -> futures::future::BoxFuture<'static, Vec<Result<SequenceHash, SequenceHash>>> {
let fail_set = self.fail_hashes.clone();
Box::pin(async move {
keys.into_iter()
.map(|h| if fail_set.contains(&h) { Err(h) } else { Ok(h) })
.collect()
})
}
fn get_blocks(
&self,
keys: Vec<SequenceHash>,
_layout: LogicalLayoutHandle,
_block_ids: Vec<BlockId>,
) -> futures::future::BoxFuture<'static, Vec<Result<SequenceHash, SequenceHash>>> {
Box::pin(async move { keys.into_iter().map(Ok).collect() })
}
}
fn test_hash(n: u64) -> SequenceHash {
SequenceHash::new(n, None, 0)
}
async fn test_phase_guard() -> BatchPhaseGuard {
let tracker = PipelineSettlementTracker::new(1);
tracker.queue_batch();
let permit = Arc::new(Semaphore::new(1))
.acquire_owned()
.await
.expect("test semaphore should remain open");
BatchPhaseGuard::starting(tracker, permit)
}
#[test]
fn block_transfer_failure_finalizes_queued_and_in_flight_handles() {
let mut transfer_states = std::collections::HashMap::new();
let mut handles = Vec::new();
for (block_id, status) in [
(7, TransferStatus::Queued),
(11, TransferStatus::Transferring),
] {
let transfer_id = TransferId::new();
let (mut state, handle) = TransferState::new(transfer_id, vec![block_id]);
state.add_passed([block_id]);
state.set_status(status);
if status == TransferStatus::Transferring {
state.mark_in_flight([block_id]);
}
let state = Arc::new(std::sync::Mutex::new(state));
transfer_states.insert(transfer_id, (state.clone(), vec![block_id]));
handles.push((block_id, state, handle));
}
let error = anyhow::anyhow!("injected block transfer failure");
BlockTransferExecutor::<crate::G1, crate::G2>::fail_transfer_states(
&transfer_states,
&error,
);
for (block_id, state, mut handle) in handles {
let result = handle
.wait()
.now_or_never()
.expect("failed transfer handle must be ready")
.expect("failed transfer must publish a result");
assert_eq!(result.status, TransferStatus::Failed);
assert_eq!(result.failed_blocks, vec![block_id]);
assert!(result.completed_blocks.is_empty());
assert_eq!(
result.error.as_deref(),
Some("block transfer failed: injected block transfer failure")
);
assert!(state.lock().unwrap().in_flight.is_empty());
}
}
#[tokio::test]
async fn test_execute_transfer_partial_failure() {
use crate::offload::handle::{TransferState, TransferStatus};
let hash_ok_1 = test_hash(1);
let hash_fail = test_hash(2);
let hash_ok_2 = test_hash(3);
let fail_hashes = [hash_fail].into_iter().collect();
let object_ops: Arc<dyn crate::object::ObjectBlockOps> =
Arc::new(FailableObjectBlockOps { fail_hashes });
let shared = SharedObjectExecutorState {
object_ops,
src_layout: LogicalLayoutHandle::G2,
skip_transfers: false,
lock_manager: None,
};
let transfer_id = crate::offload::handle::TransferId::new();
let (mut state, handle) = TransferState::new(transfer_id, vec![10, 20, 30]);
state.add_passed(vec![10, 20, 30]);
state.mark_in_flight(vec![10, 20, 30]);
let state_arc = Arc::new(std::sync::Mutex::new(state));
let blocks = vec![
ResolvedBlock::<crate::G2> {
transfer_id,
block_id: 10,
sequence_hash: hash_ok_1,
guard: None,
state: state_arc.clone(),
},
ResolvedBlock::<crate::G2> {
transfer_id,
block_id: 20,
sequence_hash: hash_fail,
guard: None,
state: state_arc.clone(),
},
ResolvedBlock::<crate::G2> {
transfer_id,
block_id: 30,
sequence_hash: hash_ok_2,
guard: None,
state: state_arc.clone(),
},
];
let mut timing = TimingTrace::new();
timing.mark_policy_complete();
timing.mark_precondition_complete();
let batch = ResolvedBatch {
blocks,
evicted: Vec::new(),
timing,
};
let mut phase = test_phase_guard().await;
ObjectTransferExecutor::<crate::G2>::execute_transfer(&shared, batch, &mut phase)
.await
.expect("execute_transfer should succeed");
phase.finish_success();
let state_guard = state_arc.lock().unwrap();
assert_eq!(handle.completed_blocks(), vec![10, 30]);
assert_eq!(handle.failed_blocks(), vec![20]);
assert_eq!(state_guard.in_flight.len(), 0);
assert_eq!(state_guard.status, TransferStatus::Failed);
assert!(state_guard.error.is_some());
drop(state_guard);
assert_eq!(handle.completed_blocks(), vec![10, 30]);
assert_eq!(handle.failed_blocks(), vec![20]);
}
#[tokio::test]
async fn test_execute_transfer_all_success() {
use crate::offload::handle::{TransferState, TransferStatus};
let hash1 = test_hash(1);
let hash2 = test_hash(2);
let object_ops: Arc<dyn crate::object::ObjectBlockOps> = Arc::new(FailableObjectBlockOps {
fail_hashes: std::collections::HashSet::new(),
});
let shared = SharedObjectExecutorState {
object_ops,
src_layout: LogicalLayoutHandle::G2,
skip_transfers: false,
lock_manager: None,
};
let transfer_id = crate::offload::handle::TransferId::new();
let (mut state, handle) = TransferState::new(transfer_id, vec![10, 20]);
state.add_passed(vec![10, 20]);
state.mark_in_flight(vec![10, 20]);
let state_arc = Arc::new(std::sync::Mutex::new(state));
let blocks = vec![
ResolvedBlock::<crate::G2> {
transfer_id,
block_id: 10,
sequence_hash: hash1,
guard: None,
state: state_arc.clone(),
},
ResolvedBlock::<crate::G2> {
transfer_id,
block_id: 20,
sequence_hash: hash2,
guard: None,
state: state_arc.clone(),
},
];
let mut timing = TimingTrace::new();
timing.mark_policy_complete();
timing.mark_precondition_complete();
let batch = ResolvedBatch {
blocks,
evicted: Vec::new(),
timing,
};
let mut phase = test_phase_guard().await;
ObjectTransferExecutor::<crate::G2>::execute_transfer(&shared, batch, &mut phase)
.await
.expect("execute_transfer should succeed");
phase.finish_success();
let state_guard = state_arc.lock().unwrap();
assert_eq!(handle.completed_blocks(), vec![10, 20]);
assert!(handle.failed_blocks().is_empty());
assert_eq!(state_guard.status, TransferStatus::Complete);
drop(state_guard);
assert_eq!(handle.completed_blocks(), vec![10, 20]);
assert!(handle.failed_blocks().is_empty());
}
#[tokio::test]
async fn test_execute_transfer_mixed_transfers() {
use crate::offload::handle::{TransferState, TransferStatus};
let hash_a1 = test_hash(10);
let hash_a2_fail = test_hash(20); let hash_b1 = test_hash(30);
let hash_b2 = test_hash(40);
let fail_hashes = [hash_a2_fail].into_iter().collect();
let object_ops: Arc<dyn crate::object::ObjectBlockOps> =
Arc::new(FailableObjectBlockOps { fail_hashes });
let shared = SharedObjectExecutorState {
object_ops,
src_layout: LogicalLayoutHandle::G2,
skip_transfers: false,
lock_manager: None,
};
let tid_a = crate::offload::handle::TransferId::new();
let (mut state_a, handle_a) = TransferState::new(tid_a, vec![100, 200]);
state_a.add_passed(vec![100, 200]);
state_a.mark_in_flight(vec![100, 200]);
let state_a_arc = Arc::new(std::sync::Mutex::new(state_a));
let tid_b = crate::offload::handle::TransferId::new();
let (mut state_b, handle_b) = TransferState::new(tid_b, vec![300, 400]);
state_b.add_passed(vec![300, 400]);
state_b.mark_in_flight(vec![300, 400]);
let state_b_arc = Arc::new(std::sync::Mutex::new(state_b));
let blocks = vec![
ResolvedBlock::<crate::G2> {
transfer_id: tid_a,
block_id: 100,
sequence_hash: hash_a1,
guard: None,
state: state_a_arc.clone(),
},
ResolvedBlock::<crate::G2> {
transfer_id: tid_a,
block_id: 200,
sequence_hash: hash_a2_fail,
guard: None,
state: state_a_arc.clone(),
},
ResolvedBlock::<crate::G2> {
transfer_id: tid_b,
block_id: 300,
sequence_hash: hash_b1,
guard: None,
state: state_b_arc.clone(),
},
ResolvedBlock::<crate::G2> {
transfer_id: tid_b,
block_id: 400,
sequence_hash: hash_b2,
guard: None,
state: state_b_arc.clone(),
},
];
let mut timing = TimingTrace::new();
timing.mark_policy_complete();
timing.mark_precondition_complete();
let batch = ResolvedBatch {
blocks,
evicted: Vec::new(),
timing,
};
let mut phase = test_phase_guard().await;
ObjectTransferExecutor::<crate::G2>::execute_transfer(&shared, batch, &mut phase)
.await
.expect("execute_transfer should succeed");
phase.finish_success();
let sa = state_a_arc.lock().unwrap();
assert_eq!(handle_a.completed_blocks(), vec![100]);
assert_eq!(handle_a.failed_blocks(), vec![200]);
assert_eq!(sa.status, TransferStatus::Failed);
assert!(sa.error.is_some());
drop(sa);
assert_eq!(handle_a.completed_blocks(), vec![100]);
assert_eq!(handle_a.failed_blocks(), vec![200]);
let sb = state_b_arc.lock().unwrap();
assert_eq!(handle_b.completed_blocks(), vec![300, 400]);
assert!(handle_b.failed_blocks().is_empty());
assert_eq!(sb.status, TransferStatus::Complete);
drop(sb);
assert_eq!(handle_b.completed_blocks(), vec![300, 400]);
assert!(handle_b.failed_blocks().is_empty());
}
}