use anyhow::Result;
use std::collections::HashMap;
use std::sync::Arc;
use crate::{
BlockId, G2, G3, InstanceId, SequenceHash,
leader::InstanceLeader,
worker::{DirectWorker, Worker},
};
use kvbm_logical::manager::BlockManager;
use kvbm_physical::manager::{LayoutHandle, TransferManager};
use kvbm_physical::transfer::StorageKind;
use kvbm_physical::{
layout::LayoutConfig,
transfer::{BlockChecksum, FillPattern},
};
use super::{managers, messenger, physical, token_blocks};
pub const DEFAULT_NUM_LAYERS: usize = 3;
pub struct TestInstanceLeader {
pub instance_id: InstanceId,
pub leader: InstanceLeader,
pub g2_manager: Arc<BlockManager<G2>>,
pub g3_manager: Option<Arc<BlockManager<G3>>>,
}
pub struct InstanceLeaderPair {
pub leader_a: TestInstanceLeader,
pub leader_b: TestInstanceLeader,
}
pub async fn create_instance_leader_pair(
block_count: usize,
block_size: usize,
) -> Result<InstanceLeaderPair> {
let messenger::MessengerPair {
messenger_a,
messenger_b,
} = messenger::create_messenger_pair_tcp().await?;
let registry_a = managers::TestRegistryBuilder::new().build();
let registry_b = managers::TestRegistryBuilder::new().build();
let g2_manager_a = Arc::new(
managers::TestManagerBuilder::<G2>::new()
.block_count(block_count)
.block_size(block_size)
.registry(registry_a.clone())
.build(),
);
let g3_manager_a = Arc::new(
managers::TestManagerBuilder::<G3>::new()
.block_count(block_count)
.block_size(block_size)
.registry(registry_a.clone())
.build(),
);
let g2_manager_b = Arc::new(
managers::TestManagerBuilder::<G2>::new()
.block_count(block_count)
.block_size(block_size)
.registry(registry_b.clone())
.build(),
);
let g3_manager_b = Arc::new(
managers::TestManagerBuilder::<G3>::new()
.block_count(block_count)
.block_size(block_size)
.registry(registry_b.clone())
.build(),
);
let leader_a = InstanceLeader::builder()
.messenger(messenger_a.clone())
.registry(registry_a.clone())
.g2_manager(g2_manager_a.clone())
.g3_manager(g3_manager_a.clone())
.workers(vec![]) .remote_leaders(vec![messenger_b.instance_id()])
.build()?;
leader_a.register_handlers()?;
let leader_b = InstanceLeader::builder()
.messenger(messenger_b.clone())
.registry(registry_b.clone())
.g2_manager(g2_manager_b.clone())
.g3_manager(g3_manager_b.clone())
.workers(vec![]) .remote_leaders(vec![messenger_a.instance_id()])
.build()?;
leader_b.register_handlers()?;
Ok(InstanceLeaderPair {
leader_a: TestInstanceLeader {
instance_id: messenger_a.instance_id(),
leader: leader_a,
g2_manager: g2_manager_a,
g3_manager: Some(g3_manager_a),
},
leader_b: TestInstanceLeader {
instance_id: messenger_b.instance_id(),
leader: leader_b,
g2_manager: g2_manager_b,
g3_manager: Some(g3_manager_b),
},
})
}
pub fn populate_leader_with_blocks(
leader: &TestInstanceLeader,
num_blocks: usize,
block_size: usize,
start_token: u32,
) -> Result<(Arc<BlockManager<G2>>, Vec<SequenceHash>)> {
let token_sequence =
super::token_blocks::create_token_sequence(num_blocks, block_size, start_token);
let seq_hashes =
managers::populate_manager_with_blocks(&leader.g2_manager, token_sequence.blocks())?;
Ok((leader.g2_manager.clone(), seq_hashes))
}
pub struct TestWorker {
pub instance_id: InstanceId,
pub worker_id: u64,
pub worker: Arc<DirectWorker>,
pub manager: Arc<TransferManager>,
pub g2_handle: LayoutHandle,
}
impl TestWorker {
pub fn fill_g2_blocks(
&self,
block_ids: &[BlockId],
pattern: FillPattern,
) -> Result<HashMap<BlockId, BlockChecksum>> {
physical::fill_and_checksum_manager(&self.manager, self.g2_handle, block_ids, pattern)
}
pub fn compute_g2_checksums(
&self,
block_ids: &[BlockId],
) -> Result<HashMap<BlockId, BlockChecksum>> {
physical::compute_manager_checksums(&self.manager, self.g2_handle, block_ids)
}
}
pub struct TestInstanceLeaderWithWorkers {
pub instance_id: InstanceId,
pub leader: InstanceLeader,
pub g2_manager: Arc<BlockManager<G2>>,
pub g3_manager: Option<Arc<BlockManager<G3>>>,
pub workers: Vec<TestWorker>,
}
impl TestInstanceLeaderWithWorkers {
pub fn g2_layout_handle(&self) -> Option<LayoutHandle> {
self.workers.first().map(|w| w.g2_handle)
}
pub fn populate_g2_blocks(
&self,
num_blocks: usize,
block_size: usize,
start_token: u32,
) -> Result<(Vec<BlockId>, Vec<SequenceHash>)> {
let token_sequence =
token_blocks::create_token_sequence(num_blocks, block_size, start_token);
let seq_hashes =
managers::populate_manager_with_blocks(&self.g2_manager, token_sequence.blocks())?;
let matched = self.g2_manager.match_blocks(&seq_hashes);
let block_ids: Vec<BlockId> = matched.into_iter().map(|b| b.block_id()).collect();
Ok((block_ids, seq_hashes))
}
pub fn fill_blocks_with_layer_pattern(
&self,
block_ids: &[BlockId],
layer: usize,
) -> Result<HashMap<BlockId, BlockChecksum>> {
let pattern = FillPattern::Constant(0xA0 + layer as u8);
let mut all_checksums = HashMap::new();
for worker in &self.workers {
let checksums = worker.fill_g2_blocks(block_ids, pattern)?;
all_checksums.extend(checksums);
}
Ok(all_checksums)
}
pub fn verify_layer_checksums(
&self,
block_ids: &[BlockId],
expected_checksums: &HashMap<BlockId, BlockChecksum>,
) -> Result<()> {
for worker in &self.workers {
let actual_checksums = worker.compute_g2_checksums(block_ids)?;
for block_id in block_ids {
let expected = expected_checksums.get(block_id).ok_or_else(|| {
anyhow::anyhow!("Missing expected checksum for block {}", block_id)
})?;
let actual = actual_checksums.get(block_id).ok_or_else(|| {
anyhow::anyhow!("Missing actual checksum for block {}", block_id)
})?;
if expected != actual {
anyhow::bail!(
"Checksum mismatch for block {}: expected {:?}, got {:?}",
block_id,
expected,
actual
);
}
}
}
Ok(())
}
}
use crate::leader::session::{
BlockInfo, EndpointSessionHandle, SessionHandle as UnifiedSessionHandle, SessionId,
SessionPhase, SessionStateSnapshot,
};
use kvbm_physical::transfer::TransferCompleteNotification;
use std::time::Duration;
pub struct TestSession {
pub session_id: SessionId,
pub endpoint_handle: EndpointSessionHandle,
pub controller_handle: UnifiedSessionHandle,
pub initial_state: SessionStateSnapshot,
}
impl TestSession {
pub async fn establish_default(
endpoint_leader: &InstanceLeader,
controller_leader: &InstanceLeader,
hashes: &[SequenceHash],
) -> Result<Self> {
Self::establish(
endpoint_leader,
controller_leader,
hashes,
Duration::from_secs(5),
)
.await
}
pub async fn establish(
endpoint_leader: &InstanceLeader,
controller_leader: &InstanceLeader,
hashes: &[SequenceHash],
timeout_duration: Duration,
) -> Result<Self> {
let (session_id, endpoint_handle) = endpoint_leader.create_endpoint_session(hashes)?;
let endpoint_instance_id = endpoint_leader.messenger().instance_id();
let mut controller_handle = controller_leader
.attach_session(endpoint_instance_id, session_id)
.await?;
let initial_state =
tokio::time::timeout(timeout_duration, controller_handle.wait_for_ready())
.await
.map_err(|_| anyhow::anyhow!("Timeout waiting for session to become ready"))?
.map_err(|e| anyhow::anyhow!("Session ready failed: {}", e))?;
Ok(Self {
session_id,
endpoint_handle,
controller_handle,
initial_state,
})
}
pub fn g2_blocks(&self) -> &[BlockInfo] {
&self.initial_state.g2_blocks
}
pub fn g3_pending(&self) -> usize {
self.initial_state.g3_pending
}
pub fn phase(&self) -> &SessionPhase {
&self.initial_state.phase
}
pub async fn pull_blocks_rdma(
&mut self,
src_blocks: &[BlockInfo],
dst_ids: &[BlockId],
) -> Result<TransferCompleteNotification> {
self.controller_handle
.pull_blocks_rdma(src_blocks, dst_ids)
.await
}
pub async fn notify_layers_ready(&self, layer_range: std::ops::Range<usize>) -> Result<()> {
self.endpoint_handle.notify_layers_ready(layer_range).await
}
pub async fn mark_blocks_pulled(&mut self, hashes: Vec<SequenceHash>) -> Result<()> {
self.controller_handle.mark_blocks_pulled(hashes).await
}
pub async fn close_endpoint(&self) -> Result<()> {
self.endpoint_handle.close().await
}
pub async fn close(self) -> Result<()> {
self.controller_handle.detach().await.ok();
self.endpoint_handle.close().await.ok();
Ok(())
}
}
pub struct InstanceLeaderPairWithWorkers {
pub decode: TestInstanceLeaderWithWorkers,
pub prefill: TestInstanceLeaderWithWorkers,
}
pub fn create_direct_worker(
instance_id: InstanceId,
agent_name: &str,
layout_config: &LayoutConfig,
storage: StorageKind,
) -> Result<TestWorker> {
let worker_id = instance_id.worker_id().as_u64();
let event_system = velo::EventManager::local();
let test_agent = physical::TestAgentBuilder::new(agent_name)
.require_backend("UCX")
.build()?;
let agent = test_agent.into_nixl_agent();
let manager = TransferManager::builder()
.event_system(Arc::new(event_system))
.nixl_agent(agent.clone())
.cuda_device_id(0)
.build()?;
let layout = physical::create_fc_layout_with_config(agent, storage, layout_config.clone());
let g2_handle = manager.register_layout(layout)?;
let direct_worker = DirectWorker::builder()
.manager(manager.clone())
.g2_handle(g2_handle)
.build()?;
Ok(TestWorker {
instance_id,
worker_id,
worker: Arc::new(direct_worker),
manager: Arc::new(manager),
g2_handle,
})
}
pub fn create_direct_workers(
num_workers: usize,
layout_config: &LayoutConfig,
storage: StorageKind,
agent_name_prefix: &str,
) -> Result<Vec<TestWorker>> {
let mut workers = Vec::with_capacity(num_workers);
for i in 0..num_workers {
let instance_id = InstanceId::new_v4();
let agent_name = format!("{}-worker-{}", agent_name_prefix, i);
let worker = create_direct_worker(instance_id, &agent_name, layout_config, storage)?;
workers.push(worker);
}
Ok(workers)
}
#[allow(clippy::too_many_arguments)]
pub async fn create_instance_leader_with_workers(
block_count: usize,
block_size: usize,
num_workers: usize,
layout_config: &LayoutConfig,
storage: StorageKind,
messenger: Arc<velo::Messenger>,
remote_leaders: Vec<InstanceId>,
agent_name_prefix: &str,
) -> Result<TestInstanceLeaderWithWorkers> {
let registry = managers::TestRegistryBuilder::new().build();
let g2_manager = Arc::new(
managers::TestManagerBuilder::<G2>::new()
.block_count(block_count)
.block_size(block_size)
.registry(registry.clone())
.build(),
);
let g3_manager = Arc::new(
managers::TestManagerBuilder::<G3>::new()
.block_count(block_count)
.block_size(block_size)
.registry(registry.clone())
.build(),
);
let workers = create_direct_workers(num_workers, layout_config, storage, agent_name_prefix)?;
let worker_refs: Vec<Arc<dyn Worker>> = workers
.iter()
.map(|w| w.worker.clone() as Arc<dyn Worker>)
.collect();
let leader = InstanceLeader::builder()
.messenger(messenger.clone())
.registry(registry.clone())
.g2_manager(g2_manager.clone())
.g3_manager(g3_manager.clone())
.workers(worker_refs)
.remote_leaders(remote_leaders)
.build()?;
leader.register_handlers()?;
Ok(TestInstanceLeaderWithWorkers {
instance_id: messenger.instance_id(),
leader,
g2_manager,
g3_manager: Some(g3_manager),
workers,
})
}
pub async fn create_instance_leader_pair_with_workers(
block_count: usize,
block_size: usize,
num_workers: usize,
layout_config: &LayoutConfig,
storage: StorageKind,
) -> Result<InstanceLeaderPairWithWorkers> {
let messenger::MessengerPair {
messenger_a,
messenger_b,
} = messenger::create_messenger_pair_tcp().await?;
let decode = create_instance_leader_with_workers(
block_count,
block_size,
num_workers,
layout_config,
storage,
messenger_a.clone(),
vec![messenger_b.instance_id()],
"decode",
)
.await?;
let prefill = create_instance_leader_with_workers(
block_count,
block_size,
num_workers,
layout_config,
storage,
messenger_b.clone(),
vec![messenger_a.instance_id()],
"prefill",
)
.await?;
Ok(InstanceLeaderPairWithWorkers { decode, prefill })
}
#[cfg(test)]
mod tests {
use super::*;
use crate::leader::ControllableSessionOptions;
#[tokio::test]
async fn test_create_instance_leader_pair() {
let pair = create_instance_leader_pair(100, 16)
.await
.expect("Should create leader pair");
assert_ne!(pair.leader_a.instance_id, pair.leader_b.instance_id);
assert_eq!(pair.leader_a.g2_manager.total_blocks(), 100);
assert_eq!(pair.leader_a.g2_manager.block_size(), 16);
assert_eq!(pair.leader_b.g2_manager.total_blocks(), 100);
assert_eq!(pair.leader_b.g2_manager.block_size(), 16);
}
#[tokio::test]
async fn test_populate_leader_with_blocks() {
let pair = create_instance_leader_pair(50, 4)
.await
.expect("Should create pair");
let (manager, hashes) =
populate_leader_with_blocks(&pair.leader_a, 10, 4, 0).expect("Should populate");
assert_eq!(hashes.len(), 10);
assert_eq!(manager.available_blocks(), 50);
let matched = manager.match_blocks(&hashes);
assert_eq!(matched.len(), 10);
}
#[tokio::test]
async fn test_scan_with_policy_linear_scan() {
use crate::leader::TieredBlock;
let pair = create_instance_leader_pair(50, 4)
.await
.expect("Should create pair");
let (_, hashes) =
populate_leader_with_blocks(&pair.leader_a, 10, 4, 0).expect("Should populate");
let blocks: Vec<TieredBlock> =
pair.leader_a
.leader
.scan_with_policy(&hashes, true, |hashes, ctx| {
for hash in hashes {
if let Some(block) = ctx.accessor().find(*hash) {
ctx.yield_item(block);
}
}
});
assert_eq!(blocks.len(), 10);
for block in &blocks {
assert!(block.is_g2());
}
for (i, block) in blocks.iter().enumerate() {
assert_eq!(block.position(), i as u64);
}
}
#[tokio::test]
async fn test_scan_with_policy_partial_matches() {
use crate::leader::TieredBlock;
let pair = create_instance_leader_pair(50, 4)
.await
.expect("Should create pair");
let (_, hashes) =
populate_leader_with_blocks(&pair.leader_a, 5, 4, 0).expect("Should populate");
let token_seq = token_blocks::create_token_sequence(3, 4, 1000);
let nonexistent_hashes = token_blocks::generate_sequence_hashes(&token_seq);
let mixed_hashes: Vec<_> = hashes
.iter()
.take(2)
.chain(nonexistent_hashes.iter().take(2))
.chain(hashes.iter().skip(2))
.copied()
.collect();
let blocks: Vec<TieredBlock> =
pair.leader_a
.leader
.scan_with_policy(&mixed_hashes, true, |hashes, ctx| {
for hash in hashes {
if let Some(block) = ctx.accessor().find(*hash) {
ctx.yield_item(block);
}
}
});
assert_eq!(blocks.len(), 5);
}
#[tokio::test]
async fn test_scan_with_policy_contiguous_single_run() {
use crate::leader::TieredBlock;
let pair = create_instance_leader_pair(100, 4)
.await
.expect("Should create pair");
let (_, hashes) =
populate_leader_with_blocks(&pair.leader_a, 10, 4, 0).expect("Should populate");
let runs: Vec<Vec<TieredBlock>> =
pair.leader_a
.leader
.scan_with_policy(&hashes, true, |hashes, ctx| {
let mut sorted_hashes = hashes.to_vec();
sorted_hashes.sort_by_key(|h| h.position());
let mut current_run = Vec::new();
let mut last_pos: Option<u64> = None;
for hash in &sorted_hashes {
if let Some(block) = ctx.accessor().find(*hash) {
let pos = block.position();
let is_contiguous = last_pos.is_none_or(|p| pos == p + 1);
if is_contiguous {
current_run.push(block);
} else {
if !current_run.is_empty() {
ctx.yield_item(std::mem::take(&mut current_run));
}
current_run.push(block);
}
last_pos = Some(pos);
} else if !current_run.is_empty() {
ctx.yield_item(std::mem::take(&mut current_run));
last_pos = None;
}
}
if !current_run.is_empty() {
ctx.yield_item(current_run);
}
});
assert_eq!(runs.len(), 1, "Expected single contiguous run");
assert_eq!(runs[0].len(), 10, "Run should contain all 10 blocks");
for (i, block) in runs[0].iter().enumerate() {
assert_eq!(block.position(), i as u64);
}
}
#[tokio::test]
async fn test_scan_with_policy_contiguous_with_gaps() {
use crate::leader::TieredBlock;
let pair = create_instance_leader_pair(100, 4)
.await
.expect("Should create pair");
let (_, all_hashes) =
populate_leader_with_blocks(&pair.leader_a, 10, 4, 0).expect("Should populate");
let query_hashes: Vec<_> = all_hashes
.iter()
.enumerate()
.filter(|(i, _)| matches!(*i, 0..=2 | 5..=6 | 8..=9))
.map(|(_, h)| *h)
.collect();
let runs: Vec<Vec<TieredBlock>> =
pair.leader_a
.leader
.scan_with_policy(&query_hashes, true, |hashes, ctx| {
let mut sorted_hashes = hashes.to_vec();
sorted_hashes.sort_by_key(|h| h.position());
let mut current_run = Vec::new();
let mut last_pos: Option<u64> = None;
for hash in &sorted_hashes {
if let Some(block) = ctx.accessor().find(*hash) {
let pos = block.position();
let is_contiguous = last_pos.is_none_or(|p| pos == p + 1);
if is_contiguous {
current_run.push(block);
} else {
if !current_run.is_empty() {
ctx.yield_item(std::mem::take(&mut current_run));
}
current_run.push(block);
}
last_pos = Some(pos);
} else if !current_run.is_empty() {
ctx.yield_item(std::mem::take(&mut current_run));
last_pos = None;
}
}
if !current_run.is_empty() {
ctx.yield_item(current_run);
}
});
assert_eq!(runs.len(), 3, "Expected 3 contiguous runs");
assert_eq!(runs[0].len(), 3);
assert_eq!(runs[0][0].position(), 0);
assert_eq!(runs[0][1].position(), 1);
assert_eq!(runs[0][2].position(), 2);
assert_eq!(runs[1].len(), 2);
assert_eq!(runs[1][0].position(), 5);
assert_eq!(runs[1][1].position(), 6);
assert_eq!(runs[2].len(), 2);
assert_eq!(runs[2][0].position(), 8);
assert_eq!(runs[2][1].position(), 9);
}
#[tokio::test]
async fn test_scan_with_policy_tiered_g2_g3() {
use crate::leader::TieredBlock;
let pair = create_instance_leader_pair(50, 4)
.await
.expect("Should create pair");
let token_sequence = token_blocks::create_token_sequence(4, 4, 0);
let all_token_blocks = token_sequence.blocks();
let g3_manager = pair
.leader_a
.g3_manager
.as_ref()
.expect("G3 manager should exist");
let g3_hashes =
managers::populate_manager_with_blocks(g3_manager, all_token_blocks).expect("G3 pop");
let even_token_blocks: Vec<_> = all_token_blocks
.iter()
.enumerate()
.filter(|(i, _)| i % 2 == 0)
.map(|(_, b)| b.clone())
.collect();
let _g2_hashes =
managers::populate_manager_with_blocks(&pair.leader_a.g2_manager, &even_token_blocks)
.expect("G2 pop");
let blocks: Vec<TieredBlock> =
pair.leader_a
.leader
.scan_with_policy(&g3_hashes, true, |hashes, ctx| {
for hash in hashes {
if let Some(block) = ctx.accessor().find(*hash) {
ctx.yield_item(block);
}
}
});
assert_eq!(blocks.len(), 4, "Should find all 4 blocks");
let g2_count = blocks.iter().filter(|b| b.is_g2()).count();
let g3_count = blocks.iter().filter(|b| b.is_g3()).count();
assert_eq!(g2_count, 2, "Should have 2 G2 blocks (even positions)");
assert_eq!(g3_count, 2, "Should have 2 G3 blocks (odd positions)");
assert!(blocks[0].is_g2(), "Block at position 0 should be G2 (even)");
assert!(blocks[1].is_g3(), "Block at position 1 should be G3 (odd)");
assert!(blocks[2].is_g2(), "Block at position 2 should be G2 (even)");
assert!(blocks[3].is_g3(), "Block at position 3 should be G3 (odd)");
for (i, block) in blocks.iter().enumerate() {
assert_eq!(
block.position(),
i as u64,
"Block {} should be at position {}",
i,
i
);
}
}
#[tokio::test]
async fn test_scan_with_policy_empty_hashes() {
use crate::leader::TieredBlock;
let pair = create_instance_leader_pair(50, 4)
.await
.expect("Should create pair");
let empty_hashes: Vec<SequenceHash> = vec![];
let blocks: Vec<TieredBlock> =
pair.leader_a
.leader
.scan_with_policy(&empty_hashes, true, |hashes, ctx| {
for hash in hashes {
if let Some(block) = ctx.accessor().find(*hash) {
ctx.yield_item(block);
}
}
});
assert!(blocks.is_empty());
}
#[tokio::test]
async fn test_scan_with_policy_yield_items() {
use crate::leader::TieredBlock;
let pair = create_instance_leader_pair(50, 4)
.await
.expect("Should create pair");
let (_, hashes) =
populate_leader_with_blocks(&pair.leader_a, 10, 4, 0).expect("Should populate");
let blocks: Vec<TieredBlock> =
pair.leader_a
.leader
.scan_with_policy(&hashes, true, |hashes, ctx| {
let found: Vec<TieredBlock> = hashes
.iter()
.filter_map(|hash| ctx.accessor().find(*hash))
.collect();
ctx.yield_items(found);
});
assert_eq!(blocks.len(), 10);
}
#[tokio::test]
async fn test_scan_with_policy_touch_parameter() {
use crate::leader::TieredBlock;
let pair = create_instance_leader_pair(50, 4)
.await
.expect("Should create pair");
let (_, hashes) =
populate_leader_with_blocks(&pair.leader_a, 5, 4, 0).expect("Should populate");
let blocks_no_touch: Vec<TieredBlock> =
pair.leader_a
.leader
.scan_with_policy(&hashes, false, |hashes, ctx| {
assert!(!ctx.accessor().touch());
for hash in hashes {
if let Some(block) = ctx.accessor().find(*hash) {
ctx.yield_item(block);
}
}
});
drop(blocks_no_touch);
let blocks_with_touch: Vec<TieredBlock> =
pair.leader_a
.leader
.scan_with_policy(&hashes, true, |hashes, ctx| {
assert!(ctx.accessor().touch());
for hash in hashes {
if let Some(block) = ctx.accessor().find(*hash) {
ctx.yield_item(block);
}
}
});
assert_eq!(blocks_with_touch.len(), 5);
}
const NUM_WORKERS: usize = 2; const LAYOUT_BLOCKS: usize = 16; const TEST_BLOCKS: usize = 4; const BLOCK_SIZE: usize = 4; const NUM_LAYERS: usize = 2; const OUTER_DIM: usize = 1; const PAGE_SIZE: usize = 4;
const INNER_DIM: usize = 64;
const DTYPE_WIDTH: usize = 2; const MANAGER_BLOCKS: usize = 16;
fn test_layout_config() -> LayoutConfig {
physical::custom_config(
LAYOUT_BLOCKS,
NUM_LAYERS,
OUTER_DIM,
PAGE_SIZE,
INNER_DIM,
DTYPE_WIDTH,
)
}
#[tokio::test(flavor = "multi_thread")]
async fn test_rdma_transfer_with_checksum_verification() {
use crate::leader::ControllableSessionOptions;
use std::time::Duration;
use tokio::time::timeout;
let layout_config = test_layout_config();
let pair = create_instance_leader_pair_with_workers(
MANAGER_BLOCKS,
BLOCK_SIZE,
NUM_WORKERS,
&layout_config,
StorageKind::Pinned,
)
.await
.expect("Should create leader pair with workers");
println!(
"\n=== RDMA Direction Test (1 block) ===\n\
Decode (source): instance={}, {} workers\n\
Prefill (dest): instance={}, {} workers",
pair.decode.instance_id,
pair.decode.workers.len(),
pair.prefill.instance_id,
pair.prefill.workers.len()
);
let src_block_ids: Vec<BlockId> = (0..TEST_BLOCKS as BlockId).collect();
let dst_block_ids: Vec<BlockId> = (TEST_BLOCKS..(TEST_BLOCKS * 2) as BlockId).collect();
println!(
"Testing {} blocks x {} workers: src={:?}, dst={:?}",
TEST_BLOCKS, NUM_WORKERS, src_block_ids, dst_block_ids
);
let mut decode_checksums_before_by_worker = Vec::new();
for (i, worker) in pair.decode.workers.iter().enumerate() {
let checksums = worker
.fill_g2_blocks(&src_block_ids, FillPattern::Constant(0xAA))
.expect("Should fill Decode G2 blocks");
println!(
"BEFORE transfer - Decode worker {} blocks: {:?}",
i, src_block_ids
);
decode_checksums_before_by_worker.push(checksums);
}
let mut prefill_checksums_before_by_worker = Vec::new();
for (i, worker) in pair.prefill.workers.iter().enumerate() {
let checksums = worker
.fill_g2_blocks(&dst_block_ids, FillPattern::Constant(0xBB))
.expect("Should fill Prefill G2 blocks");
println!(
"BEFORE transfer - Prefill worker {} blocks: {:?}",
i, dst_block_ids
);
prefill_checksums_before_by_worker.push(checksums);
}
assert_ne!(
decode_checksums_before_by_worker[0][&0],
prefill_checksums_before_by_worker[0][&dst_block_ids[0]],
"Pre-transfer: Decode and Prefill should have different data"
);
let test_leader = TestInstanceLeader {
instance_id: pair.decode.instance_id,
leader: pair.decode.leader.clone(),
g2_manager: pair.decode.g2_manager.clone(),
g3_manager: pair.decode.g3_manager.clone(),
};
let (_, sequence_hashes) =
populate_leader_with_blocks(&test_leader, TEST_BLOCKS, BLOCK_SIZE, 0)
.expect("Should populate leader");
let session_result = pair
.decode
.leader
.create_controllable_session_with_options(
&sequence_hashes,
ControllableSessionOptions { auto_stage: false },
)
.expect("Should create controllable session");
println!(
"Decode session created: {} G2 blocks",
session_result.local_g2_count
);
let mut handle = pair
.prefill
.leader
.attach_session(pair.decode.instance_id, session_result.session_id)
.await
.expect("Should attach");
let state = timeout(Duration::from_secs(5), handle.wait_for_ready())
.await
.expect("Timeout")
.expect("Should get state");
println!(
"Prefill sees {} G2 blocks from Decode",
state.g2_blocks.len()
);
println!("\n--- Executing RDMA pull: Decode block 0 -> Prefill block 1 ---");
let notification = handle
.pull_blocks_rdma(&state.g2_blocks, &dst_block_ids)
.await
.expect("Should initiate RDMA pull");
notification.await.expect("Transfer should complete");
println!("Transfer complete!\n");
println!("\nVerifying SPMD replication - all workers have all blocks:");
println!(
" Each worker: src={:?} -> dst={:?}",
src_block_ids, dst_block_ids
);
for (worker_idx, (decode_worker, prefill_worker)) in pair
.decode
.workers
.iter()
.zip(pair.prefill.workers.iter())
.enumerate()
{
let decode_checksums_after = decode_worker
.compute_g2_checksums(&src_block_ids)
.expect("compute Decode checksums");
let prefill_checksums_after = prefill_worker
.compute_g2_checksums(&dst_block_ids)
.expect("compute Prefill checksums");
println!(
"\nWorker {} verification ({} blocks):",
worker_idx, TEST_BLOCKS
);
let decode_checksums_before = &decode_checksums_before_by_worker[worker_idx];
for i in 0..TEST_BLOCKS {
let src_id = src_block_ids[i];
let dst_id = dst_block_ids[i];
assert_eq!(
decode_checksums_before[&src_id], decode_checksums_after[&src_id],
"Decode block {} was modified!",
src_id
);
assert_eq!(
decode_checksums_before[&src_id], prefill_checksums_after[&dst_id],
"Prefill block {} doesn't have Decode block {}'s data",
dst_id, src_id
);
println!(" Worker {} block {} -> {}", worker_idx, src_id, dst_id);
}
}
println!(
"\n=== SUCCESS: {} blocks correctly transferred across {} workers (SPMD) ===",
TEST_BLOCKS, NUM_WORKERS
);
handle.mark_blocks_pulled(sequence_hashes).await.ok();
handle.detach().await.ok();
}
#[tokio::test(flavor = "multi_thread")]
async fn test_bidirectional_layerwise_transfer() {
use crate::leader::session::SessionHandle as UnifiedSessionHandle;
use std::time::Duration;
use tokio::time::timeout;
const CACHED_BLOCKS: usize = 4; const NEW_BLOCKS: usize = 2; const NUM_TEST_LAYERS: usize = NUM_LAYERS;
let layout_config = test_layout_config();
let pair = create_instance_leader_pair_with_workers(
MANAGER_BLOCKS,
BLOCK_SIZE,
NUM_WORKERS,
&layout_config,
StorageKind::Pinned,
)
.await
.expect("Should create leader pair with workers");
println!(
"\n=== Bidirectional Layerwise Transfer Test ===\n\
Decode: instance={}, {} workers\n\
Prefill: instance={}, {} workers\n\
Cached blocks: {}, New blocks: {}, Layers: {}",
pair.decode.instance_id,
pair.decode.workers.len(),
pair.prefill.instance_id,
pair.prefill.workers.len(),
CACHED_BLOCKS,
NEW_BLOCKS,
NUM_TEST_LAYERS
);
println!("\n--- Phase 1: Decode Setup ---");
let (decode_cached_block_ids, cached_hashes) = pair
.decode
.populate_g2_blocks(CACHED_BLOCKS, BLOCK_SIZE, 0)
.expect("Should populate Decode");
for worker in &pair.decode.workers {
worker
.fill_g2_blocks(&decode_cached_block_ids, FillPattern::Constant(0xCA))
.expect("Should fill cached blocks");
}
println!(
"Decode populated with {} cached blocks: {:?}",
CACHED_BLOCKS, decode_cached_block_ids
);
let (prefill_new_block_ids, new_hashes) = pair
.prefill
.populate_g2_blocks(NEW_BLOCKS, BLOCK_SIZE, 1000)
.expect("Should populate Prefill");
println!(
"Prefill populated with {} new blocks: {:?}",
NEW_BLOCKS, prefill_new_block_ids
);
println!("\n--- Phase 2: Prefill Pulls from Decode ---");
let session_result = pair
.decode
.leader
.create_controllable_session_with_options(
&cached_hashes,
ControllableSessionOptions { auto_stage: false },
)
.expect("Should create controllable session");
println!(
"Decode session created: {} G2 blocks",
session_result.local_g2_count
);
let mut prefill_handle = pair
.prefill
.leader
.attach_session(pair.decode.instance_id, session_result.session_id)
.await
.expect("Should attach");
let state = timeout(Duration::from_secs(5), prefill_handle.wait_for_ready())
.await
.expect("Timeout waiting for initial state")
.expect("Should get initial state");
println!(
"Prefill sees {} G2 blocks from Decode",
state.g2_blocks.len()
);
let prefill_dst_blocks = pair
.prefill
.g2_manager
.allocate_blocks(CACHED_BLOCKS)
.expect("Should allocate destination blocks on Prefill");
let prefill_dst_block_ids: Vec<BlockId> =
prefill_dst_blocks.iter().map(|b| b.block_id()).collect();
println!(
"Prefill allocated destination blocks: {:?}",
prefill_dst_block_ids
);
let notification = prefill_handle
.pull_blocks_rdma(&state.g2_blocks, &prefill_dst_block_ids)
.await
.expect("Should initiate RDMA pull");
notification.await.expect("Transfer should complete");
println!("Prefill pulled {} cached blocks", CACHED_BLOCKS);
println!("Verifying Prefill received Decode's cached data...");
for (worker_idx, (decode_worker, prefill_worker)) in pair
.decode
.workers
.iter()
.zip(pair.prefill.workers.iter())
.enumerate()
{
let decode_checksums = decode_worker
.compute_g2_checksums(&decode_cached_block_ids)
.expect("Should compute Decode checksums");
let prefill_checksums = prefill_worker
.compute_g2_checksums(&prefill_dst_block_ids)
.expect("Should compute Prefill checksums");
for i in 0..CACHED_BLOCKS {
let src_id = decode_cached_block_ids[i];
let dst_id = prefill_dst_block_ids[i];
assert_eq!(
decode_checksums[&src_id], prefill_checksums[&dst_id],
"Worker {}: Prefill block {} should match Decode block {}",
worker_idx, dst_id, src_id
);
}
println!(
" Worker {} verified: Prefill has Decode's cached data",
worker_idx
);
}
prefill_handle.detach().await.ok();
println!("Prefill detached (Decode keeps cached blocks)");
println!("\n--- Phase 3: Role Reversal ---");
let (prefill_session_id, prefill_session_handle) = pair
.prefill
.leader
.create_endpoint_session(&new_hashes)
.expect("Should create endpoint session");
println!("Prefill created endpoint session: {}", prefill_session_id);
let mut decode_handle: UnifiedSessionHandle = pair
.decode
.leader
.attach_session(pair.prefill.instance_id, prefill_session_id)
.await
.expect("Should attach to Prefill's session");
println!("Decode attached to Prefill's session");
let state = timeout(Duration::from_secs(5), decode_handle.wait_for_ready())
.await
.expect("Timeout waiting for ready")
.expect("Should get ready state");
println!(
"Decode sees {} G2 blocks from Prefill, phase: {:?}",
state.g2_blocks.len(),
state.phase
);
println!("\n--- Phase 4: Layerwise Transfer ---");
let decode_dst_blocks = pair
.decode
.g2_manager
.allocate_blocks(NEW_BLOCKS)
.expect("Should allocate destination blocks on Decode");
let decode_dst_block_ids: Vec<BlockId> =
decode_dst_blocks.iter().map(|b| b.block_id()).collect();
println!(
"Decode allocated destination blocks: {:?}",
decode_dst_block_ids
);
for worker in &pair.prefill.workers {
worker
.fill_g2_blocks(&prefill_new_block_ids, FillPattern::Constant(0xBB))
.expect("Should fill Prefill blocks");
}
println!("Prefill blocks filled with pattern 0xBB");
for layer in 0..NUM_TEST_LAYERS {
println!("\n Layer {} notification:", layer);
prefill_session_handle
.notify_layers_ready(0..layer + 1)
.await
.expect("Should notify layer ready");
println!(" Prefill notified layers 0..{} ready", layer + 1);
tokio::time::sleep(Duration::from_millis(10)).await;
}
println!("\n Decode pulling all layers...");
let notification = decode_handle
.pull_blocks_rdma(&state.g2_blocks, &decode_dst_block_ids)
.await
.expect("Should initiate RDMA pull");
notification.await.expect("Transfer should complete");
println!(" Decode pulled all {} blocks", NEW_BLOCKS);
println!("Verifying Decode received Prefill's data...");
for (worker_idx, (prefill_worker, decode_worker)) in pair
.prefill
.workers
.iter()
.zip(pair.decode.workers.iter())
.enumerate()
{
let prefill_checksums = prefill_worker
.compute_g2_checksums(&prefill_new_block_ids)
.expect("Should compute Prefill checksums");
let decode_checksums = decode_worker
.compute_g2_checksums(&decode_dst_block_ids)
.expect("Should compute Decode checksums");
for i in 0..NEW_BLOCKS {
let src_id = prefill_new_block_ids[i];
let dst_id = decode_dst_block_ids[i];
assert_eq!(
prefill_checksums[&src_id], decode_checksums[&dst_id],
"Worker {}: Decode block {} should match Prefill block {}",
worker_idx, dst_id, src_id
);
}
println!(
" Worker {} verified: Decode has Prefill's data (pattern 0xBB)",
worker_idx
);
}
println!("\n--- Phase 5: Cleanup ---");
decode_handle
.mark_blocks_pulled(new_hashes.clone())
.await
.ok();
decode_handle.detach().await.ok();
println!("Decode detached from Prefill's session");
prefill_session_handle.close().await.ok();
println!("Prefill closed endpoint session");
println!(
"\n=== SUCCESS: Bidirectional layerwise transfer completed ===\n\
- {} cached blocks transferred Decode -> Prefill\n\
- {} new blocks transferred Prefill -> Decode ({} layers)",
CACHED_BLOCKS, NEW_BLOCKS, NUM_TEST_LAYERS
);
}
}