use anyhow::Result;
use std::sync::Arc;
use tokio::sync::watch;
use crate::worker::group::ParallelWorkers;
use crate::{BlockId, InstanceId, SequenceHash};
use kvbm_common::LogicalLayoutHandle;
use kvbm_physical::transfer::{TransferCompleteNotification, TransferOptions};
use super::{
BlockInfo, ControlRole, SessionId, SessionMessage, SessionPhase, SessionStateSnapshot,
transport::MessageTransport,
};
pub struct SessionHandle {
session_id: SessionId,
remote_instance: InstanceId,
local_instance: InstanceId,
transport: Arc<MessageTransport>,
state_rx: watch::Receiver<SessionStateSnapshot>,
parallel_worker: Option<Arc<dyn ParallelWorkers>>,
}
impl SessionHandle {
#[allow(dead_code)]
pub(crate) fn new(
session_id: SessionId,
remote_instance: InstanceId,
local_instance: InstanceId,
transport: Arc<MessageTransport>,
state_rx: watch::Receiver<SessionStateSnapshot>,
) -> Self {
Self {
session_id,
remote_instance,
local_instance,
transport,
state_rx,
parallel_worker: None,
}
}
pub fn with_rdma_support(mut self, parallel_worker: Arc<dyn ParallelWorkers>) -> Self {
self.parallel_worker = Some(parallel_worker);
self
}
pub fn session_id(&self) -> SessionId {
self.session_id
}
pub fn remote_instance(&self) -> InstanceId {
self.remote_instance
}
pub fn local_instance(&self) -> InstanceId {
self.local_instance
}
pub fn current_state(&self) -> SessionStateSnapshot {
self.state_rx.borrow().clone()
}
pub fn phase(&self) -> SessionPhase {
self.state_rx.borrow().phase
}
pub fn remote_control_role(&self) -> ControlRole {
self.state_rx.borrow().control_role
}
pub fn has_changed(&self) -> bool {
self.state_rx.has_changed().unwrap_or(false)
}
pub async fn wait_for_change(&mut self) -> Result<SessionStateSnapshot> {
self.state_rx
.changed()
.await
.map_err(|e| anyhow::anyhow!("State channel closed: {}", e))?;
Ok(self.state_rx.borrow().clone())
}
pub async fn wait_for_ready(&mut self) -> Result<SessionStateSnapshot> {
self.state_rx
.wait_for(|s| s.phase == SessionPhase::Ready || s.phase.is_terminal())
.await
.map_err(|e| anyhow::anyhow!("Failed waiting for ready: {}", e))?;
let state = self.state_rx.borrow().clone();
if state.phase == SessionPhase::Failed {
anyhow::bail!("Session failed while waiting for ready");
}
Ok(state)
}
pub async fn wait_for_complete(&mut self) -> Result<SessionStateSnapshot> {
self.state_rx
.wait_for(|s| s.phase.is_terminal())
.await
.map_err(|e| anyhow::anyhow!("Failed waiting for complete: {}", e))?;
Ok(self.state_rx.borrow().clone())
}
pub fn is_complete(&self) -> bool {
self.state_rx.borrow().phase.is_terminal()
}
pub fn is_ready(&self) -> bool {
self.state_rx.borrow().phase == SessionPhase::Ready
}
pub fn get_g2_blocks(&self) -> Vec<BlockInfo> {
self.state_rx.borrow().g2_blocks.clone()
}
pub fn g3_pending_count(&self) -> usize {
self.state_rx.borrow().g3_pending
}
pub fn ready_layer_range(&self) -> Option<std::ops::Range<usize>> {
self.state_rx.borrow().ready_layer_range.clone()
}
pub async fn trigger_staging(&self) -> Result<()> {
let msg = SessionMessage::TriggerStaging {
session_id: self.session_id,
};
self.transport.send_session(self.remote_instance, msg).await
}
pub async fn mark_blocks_pulled(&self, pulled_hashes: Vec<SequenceHash>) -> Result<()> {
let msg = SessionMessage::BlocksPulled {
session_id: self.session_id,
pulled_hashes,
};
self.transport.send_session(self.remote_instance, msg).await
}
pub async fn detach(self) -> Result<()> {
let msg = SessionMessage::Detach {
peer: self.local_instance,
session_id: self.session_id,
};
self.transport.send_session(self.remote_instance, msg).await
}
pub async fn yield_control(&self) -> Result<()> {
let msg = SessionMessage::YieldControl {
peer: self.local_instance,
session_id: self.session_id,
};
self.transport.send_session(self.remote_instance, msg).await
}
pub async fn acquire_control(&self) -> Result<()> {
let msg = SessionMessage::AcquireControl {
peer: self.local_instance,
session_id: self.session_id,
};
self.transport.send_session(self.remote_instance, msg).await
}
pub fn has_remote_metadata(&self) -> bool {
self.parallel_worker
.as_ref()
.map(|pw| pw.has_remote_metadata(self.remote_instance))
.unwrap_or(false)
}
pub async fn ensure_metadata_imported(&mut self) -> Result<()> {
let parallel_worker = self
.parallel_worker
.as_ref()
.ok_or_else(|| anyhow::anyhow!("RDMA support not configured"))?;
if parallel_worker.has_remote_metadata(self.remote_instance) {
return Ok(());
}
let remote_metadata = self
.transport
.request_metadata(self.remote_instance)
.await?;
parallel_worker
.connect_remote(self.remote_instance, remote_metadata)?
.await?;
Ok(())
}
pub async fn pull_blocks_rdma(
&mut self,
blocks: &[BlockInfo],
local_dst_block_ids: &[BlockId],
) -> Result<TransferCompleteNotification> {
self.ensure_metadata_imported().await?;
self.pull_blocks_rdma_explicit(blocks, local_dst_block_ids)
}
pub fn pull_blocks_rdma_explicit(
&self,
blocks: &[BlockInfo],
local_dst_block_ids: &[BlockId],
) -> Result<TransferCompleteNotification> {
let parallel_worker = self
.parallel_worker
.as_ref()
.ok_or_else(|| anyhow::anyhow!("RDMA support not configured"))?;
if !parallel_worker.has_remote_metadata(self.remote_instance) {
anyhow::bail!(
"Remote metadata not imported for instance {}",
self.remote_instance
);
}
if blocks.len() != local_dst_block_ids.len() {
anyhow::bail!(
"Block count mismatch: source={}, destination={}",
blocks.len(),
local_dst_block_ids.len()
);
}
let src_block_ids: Vec<BlockId> = blocks.iter().map(|b| b.block_id).collect();
parallel_worker.execute_remote_onboard_for_instance(
self.remote_instance,
LogicalLayoutHandle::G2,
src_block_ids,
LogicalLayoutHandle::G2,
local_dst_block_ids.to_vec().into(),
Default::default(),
)
}
pub async fn pull_blocks_rdma_with_options(
&mut self,
blocks: &[BlockInfo],
local_dst_block_ids: &[BlockId],
options: TransferOptions,
) -> Result<TransferCompleteNotification> {
self.ensure_metadata_imported().await?;
self.pull_blocks_rdma_with_options_explicit(blocks, local_dst_block_ids, options)
}
pub fn pull_blocks_rdma_with_options_explicit(
&self,
blocks: &[BlockInfo],
local_dst_block_ids: &[BlockId],
options: TransferOptions,
) -> Result<TransferCompleteNotification> {
let parallel_worker = self
.parallel_worker
.as_ref()
.ok_or_else(|| anyhow::anyhow!("RDMA support not configured"))?;
if !parallel_worker.has_remote_metadata(self.remote_instance) {
anyhow::bail!(
"Remote metadata not imported for instance {}",
self.remote_instance
);
}
if blocks.len() != local_dst_block_ids.len() {
anyhow::bail!(
"Block count mismatch: source={}, destination={}",
blocks.len(),
local_dst_block_ids.len()
);
}
let src_block_ids: Vec<BlockId> = blocks.iter().map(|b| b.block_id).collect();
parallel_worker.execute_remote_onboard_for_instance(
self.remote_instance,
LogicalLayoutHandle::G2,
src_block_ids,
LogicalLayoutHandle::G2,
local_dst_block_ids.to_vec().into(),
options,
)
}
}
pub struct SessionHandleStateTx {
tx: watch::Sender<SessionStateSnapshot>,
}
impl SessionHandleStateTx {
pub fn new(tx: watch::Sender<SessionStateSnapshot>) -> Self {
Self { tx }
}
pub fn update(&self, state: SessionStateSnapshot) {
let _ = self.tx.send(state);
}
pub fn set_phase(&self, phase: SessionPhase) {
self.tx.send_modify(|state| {
state.phase = phase;
});
}
pub fn set_g2_blocks(&self, blocks: Vec<BlockInfo>) {
self.tx.send_modify(|state| {
state.g2_blocks = blocks;
});
}
pub fn add_staged_blocks(
&self,
staged: Vec<BlockInfo>,
g3_remaining: usize,
layer_range: Option<std::ops::Range<usize>>,
) {
self.tx.send_modify(|state| {
state.g2_blocks.extend(staged);
state.g3_pending = g3_remaining;
state.ready_layer_range = layer_range;
if g3_remaining == 0 && state.ready_layer_range.is_none() {
state.phase = SessionPhase::Ready;
}
});
}
pub fn set_failed(&self) {
self.tx.send_modify(|state| {
state.phase = SessionPhase::Failed;
});
}
}
pub fn session_handle_state_channel()
-> (SessionHandleStateTx, watch::Receiver<SessionStateSnapshot>) {
let initial = SessionStateSnapshot {
phase: SessionPhase::Searching,
control_role: ControlRole::Controllee,
g2_blocks: Vec::new(),
g3_pending: 0,
ready_layer_range: None,
};
let (tx, rx) = watch::channel(initial);
(SessionHandleStateTx::new(tx), rx)
}
#[cfg(test)]
mod tests {
use super::*;
use dashmap::DashMap;
fn create_test_transport() -> Arc<MessageTransport> {
Arc::new(MessageTransport::local(
Arc::new(DashMap::new()),
Arc::new(DashMap::new()),
))
}
#[test]
fn test_session_handle_state_channel() {
let (tx, rx) = session_handle_state_channel();
let state = rx.borrow().clone();
assert_eq!(state.phase, SessionPhase::Searching);
assert_eq!(state.control_role, ControlRole::Controllee);
assert!(state.g2_blocks.is_empty());
tx.set_phase(SessionPhase::Ready);
let state = rx.borrow().clone();
assert_eq!(state.phase, SessionPhase::Ready);
}
#[test]
fn test_session_handle_creation() {
let (_, rx) = session_handle_state_channel();
let transport = create_test_transport();
let session_id = SessionId::new_v4();
let remote_id = InstanceId::new_v4();
let local_id = InstanceId::new_v4();
let handle = SessionHandle::new(session_id, remote_id, local_id, transport, rx);
assert_eq!(handle.session_id(), session_id);
assert_eq!(handle.remote_instance(), remote_id);
assert_eq!(handle.local_instance(), local_id);
assert_eq!(handle.phase(), SessionPhase::Searching);
assert!(!handle.is_ready());
assert!(!handle.is_complete());
assert!(!handle.has_remote_metadata());
}
#[tokio::test]
async fn test_wait_for_ready() {
let (tx, rx) = session_handle_state_channel();
let transport = create_test_transport();
let session_id = SessionId::new_v4();
let mut handle = SessionHandle::new(
session_id,
InstanceId::new_v4(),
InstanceId::new_v4(),
transport,
rx,
);
tokio::spawn(async move {
tokio::time::sleep(tokio::time::Duration::from_millis(10)).await;
tx.set_phase(SessionPhase::Ready);
});
let state = handle.wait_for_ready().await.unwrap();
assert_eq!(state.phase, SessionPhase::Ready);
}
#[test]
fn test_add_staged_blocks() {
let (tx, rx) = session_handle_state_channel();
tx.update(SessionStateSnapshot {
phase: SessionPhase::Staging,
control_role: ControlRole::Controllee,
g2_blocks: Vec::new(),
g3_pending: 5,
ready_layer_range: None,
});
let state = rx.borrow().clone();
assert_eq!(state.g3_pending, 5);
assert!(state.g2_blocks.is_empty());
let block = BlockInfo {
block_id: 42,
sequence_hash: crate::SequenceHash::new(1, None, 100),
layout_handle: kvbm_physical::manager::LayoutHandle::new(0, 1),
};
tx.add_staged_blocks(vec![block], 0, None);
let state = rx.borrow().clone();
assert_eq!(state.g2_blocks.len(), 1);
assert_eq!(state.g3_pending, 0);
assert_eq!(state.phase, SessionPhase::Ready);
}
#[test]
fn test_set_failed() {
let (tx, rx) = session_handle_state_channel();
assert_eq!(rx.borrow().phase, SessionPhase::Searching);
tx.set_failed();
assert_eq!(rx.borrow().phase, SessionPhase::Failed);
}
#[tokio::test]
async fn test_wait_for_complete() {
let (tx, rx) = session_handle_state_channel();
let transport = create_test_transport();
let session_id = SessionId::new_v4();
let mut handle = SessionHandle::new(
session_id,
InstanceId::new_v4(),
InstanceId::new_v4(),
transport,
rx,
);
tokio::spawn(async move {
tokio::time::sleep(tokio::time::Duration::from_millis(10)).await;
tx.set_phase(SessionPhase::Complete);
});
let state = handle.wait_for_complete().await.unwrap();
assert_eq!(state.phase, SessionPhase::Complete);
assert!(handle.is_complete());
}
}