kvbm_engine/leader/session/
handle.rs1use anyhow::Result;
16use std::sync::Arc;
17use tokio::sync::watch;
18
19use crate::worker::group::ParallelWorkers;
20use crate::{BlockId, InstanceId, SequenceHash};
21use kvbm_common::LogicalLayoutHandle;
22use kvbm_physical::transfer::{TransferCompleteNotification, TransferOptions};
23
24use super::{
25 BlockInfo, ControlRole, SessionId, SessionMessage, SessionPhase, SessionStateSnapshot,
26 transport::MessageTransport,
27};
28
29pub struct SessionHandle {
61 session_id: SessionId,
62 remote_instance: InstanceId,
63 local_instance: InstanceId,
64 transport: Arc<MessageTransport>,
65
66 state_rx: watch::Receiver<SessionStateSnapshot>,
68
69 parallel_worker: Option<Arc<dyn ParallelWorkers>>,
71}
72
73impl SessionHandle {
74 #[allow(dead_code)]
79 pub(crate) fn new(
80 session_id: SessionId,
81 remote_instance: InstanceId,
82 local_instance: InstanceId,
83 transport: Arc<MessageTransport>,
84 state_rx: watch::Receiver<SessionStateSnapshot>,
85 ) -> Self {
86 Self {
87 session_id,
88 remote_instance,
89 local_instance,
90 transport,
91 state_rx,
92 parallel_worker: None,
93 }
94 }
95
96 pub fn with_rdma_support(mut self, parallel_worker: Arc<dyn ParallelWorkers>) -> Self {
98 self.parallel_worker = Some(parallel_worker);
99 self
100 }
101
102 pub fn session_id(&self) -> SessionId {
108 self.session_id
109 }
110
111 pub fn remote_instance(&self) -> InstanceId {
113 self.remote_instance
114 }
115
116 pub fn local_instance(&self) -> InstanceId {
118 self.local_instance
119 }
120
121 pub fn current_state(&self) -> SessionStateSnapshot {
127 self.state_rx.borrow().clone()
128 }
129
130 pub fn phase(&self) -> SessionPhase {
132 self.state_rx.borrow().phase
133 }
134
135 pub fn remote_control_role(&self) -> ControlRole {
137 self.state_rx.borrow().control_role
138 }
139
140 pub fn has_changed(&self) -> bool {
142 self.state_rx.has_changed().unwrap_or(false)
143 }
144
145 pub async fn wait_for_change(&mut self) -> Result<SessionStateSnapshot> {
147 self.state_rx
148 .changed()
149 .await
150 .map_err(|e| anyhow::anyhow!("State channel closed: {}", e))?;
151 Ok(self.state_rx.borrow().clone())
152 }
153
154 pub async fn wait_for_ready(&mut self) -> Result<SessionStateSnapshot> {
156 self.state_rx
157 .wait_for(|s| s.phase == SessionPhase::Ready || s.phase.is_terminal())
158 .await
159 .map_err(|e| anyhow::anyhow!("Failed waiting for ready: {}", e))?;
160
161 let state = self.state_rx.borrow().clone();
162 if state.phase == SessionPhase::Failed {
163 anyhow::bail!("Session failed while waiting for ready");
164 }
165 Ok(state)
166 }
167
168 pub async fn wait_for_complete(&mut self) -> Result<SessionStateSnapshot> {
170 self.state_rx
171 .wait_for(|s| s.phase.is_terminal())
172 .await
173 .map_err(|e| anyhow::anyhow!("Failed waiting for complete: {}", e))?;
174 Ok(self.state_rx.borrow().clone())
175 }
176
177 pub fn is_complete(&self) -> bool {
179 self.state_rx.borrow().phase.is_terminal()
180 }
181
182 pub fn is_ready(&self) -> bool {
184 self.state_rx.borrow().phase == SessionPhase::Ready
185 }
186
187 pub fn get_g2_blocks(&self) -> Vec<BlockInfo> {
189 self.state_rx.borrow().g2_blocks.clone()
190 }
191
192 pub fn g3_pending_count(&self) -> usize {
194 self.state_rx.borrow().g3_pending
195 }
196
197 pub fn ready_layer_range(&self) -> Option<std::ops::Range<usize>> {
202 self.state_rx.borrow().ready_layer_range.clone()
203 }
204
205 pub async fn trigger_staging(&self) -> Result<()> {
213 let msg = SessionMessage::TriggerStaging {
214 session_id: self.session_id,
215 };
216 self.transport.send_session(self.remote_instance, msg).await
217 }
218
219 pub async fn mark_blocks_pulled(&self, pulled_hashes: Vec<SequenceHash>) -> Result<()> {
223 let msg = SessionMessage::BlocksPulled {
224 session_id: self.session_id,
225 pulled_hashes,
226 };
227 self.transport.send_session(self.remote_instance, msg).await
228 }
229
230 pub async fn detach(self) -> Result<()> {
234 let msg = SessionMessage::Detach {
235 peer: self.local_instance,
236 session_id: self.session_id,
237 };
238 self.transport.send_session(self.remote_instance, msg).await
239 }
240
241 pub async fn yield_control(&self) -> Result<()> {
250 let msg = SessionMessage::YieldControl {
251 peer: self.local_instance,
252 session_id: self.session_id,
253 };
254 self.transport.send_session(self.remote_instance, msg).await
255 }
256
257 pub async fn acquire_control(&self) -> Result<()> {
261 let msg = SessionMessage::AcquireControl {
262 peer: self.local_instance,
263 session_id: self.session_id,
264 };
265 self.transport.send_session(self.remote_instance, msg).await
266 }
267
268 pub fn has_remote_metadata(&self) -> bool {
274 self.parallel_worker
275 .as_ref()
276 .map(|pw| pw.has_remote_metadata(self.remote_instance))
277 .unwrap_or(false)
278 }
279
280 pub async fn ensure_metadata_imported(&mut self) -> Result<()> {
282 let parallel_worker = self
283 .parallel_worker
284 .as_ref()
285 .ok_or_else(|| anyhow::anyhow!("RDMA support not configured"))?;
286
287 if parallel_worker.has_remote_metadata(self.remote_instance) {
288 return Ok(());
289 }
290
291 let remote_metadata = self
292 .transport
293 .request_metadata(self.remote_instance)
294 .await?;
295
296 parallel_worker
297 .connect_remote(self.remote_instance, remote_metadata)?
298 .await?;
299
300 Ok(())
301 }
302
303 pub async fn pull_blocks_rdma(
310 &mut self,
311 blocks: &[BlockInfo],
312 local_dst_block_ids: &[BlockId],
313 ) -> Result<TransferCompleteNotification> {
314 self.ensure_metadata_imported().await?;
315 self.pull_blocks_rdma_explicit(blocks, local_dst_block_ids)
316 }
317
318 pub fn pull_blocks_rdma_explicit(
322 &self,
323 blocks: &[BlockInfo],
324 local_dst_block_ids: &[BlockId],
325 ) -> Result<TransferCompleteNotification> {
326 let parallel_worker = self
327 .parallel_worker
328 .as_ref()
329 .ok_or_else(|| anyhow::anyhow!("RDMA support not configured"))?;
330
331 if !parallel_worker.has_remote_metadata(self.remote_instance) {
332 anyhow::bail!(
333 "Remote metadata not imported for instance {}",
334 self.remote_instance
335 );
336 }
337
338 if blocks.len() != local_dst_block_ids.len() {
339 anyhow::bail!(
340 "Block count mismatch: source={}, destination={}",
341 blocks.len(),
342 local_dst_block_ids.len()
343 );
344 }
345
346 let src_block_ids: Vec<BlockId> = blocks.iter().map(|b| b.block_id).collect();
347
348 parallel_worker.execute_remote_onboard_for_instance(
349 self.remote_instance,
350 LogicalLayoutHandle::G2,
351 src_block_ids,
352 LogicalLayoutHandle::G2,
353 local_dst_block_ids.to_vec().into(),
354 Default::default(),
355 )
356 }
357
358 pub async fn pull_blocks_rdma_with_options(
374 &mut self,
375 blocks: &[BlockInfo],
376 local_dst_block_ids: &[BlockId],
377 options: TransferOptions,
378 ) -> Result<TransferCompleteNotification> {
379 self.ensure_metadata_imported().await?;
380 self.pull_blocks_rdma_with_options_explicit(blocks, local_dst_block_ids, options)
381 }
382
383 pub fn pull_blocks_rdma_with_options_explicit(
387 &self,
388 blocks: &[BlockInfo],
389 local_dst_block_ids: &[BlockId],
390 options: TransferOptions,
391 ) -> Result<TransferCompleteNotification> {
392 let parallel_worker = self
393 .parallel_worker
394 .as_ref()
395 .ok_or_else(|| anyhow::anyhow!("RDMA support not configured"))?;
396
397 if !parallel_worker.has_remote_metadata(self.remote_instance) {
398 anyhow::bail!(
399 "Remote metadata not imported for instance {}",
400 self.remote_instance
401 );
402 }
403
404 if blocks.len() != local_dst_block_ids.len() {
405 anyhow::bail!(
406 "Block count mismatch: source={}, destination={}",
407 blocks.len(),
408 local_dst_block_ids.len()
409 );
410 }
411
412 let src_block_ids: Vec<BlockId> = blocks.iter().map(|b| b.block_id).collect();
413
414 parallel_worker.execute_remote_onboard_for_instance(
415 self.remote_instance,
416 LogicalLayoutHandle::G2,
417 src_block_ids,
418 LogicalLayoutHandle::G2,
419 local_dst_block_ids.to_vec().into(),
420 options,
421 )
422 }
423}
424
425pub struct SessionHandleStateTx {
427 tx: watch::Sender<SessionStateSnapshot>,
428}
429
430impl SessionHandleStateTx {
431 pub fn new(tx: watch::Sender<SessionStateSnapshot>) -> Self {
433 Self { tx }
434 }
435
436 pub fn update(&self, state: SessionStateSnapshot) {
438 let _ = self.tx.send(state);
439 }
440
441 pub fn set_phase(&self, phase: SessionPhase) {
443 self.tx.send_modify(|state| {
444 state.phase = phase;
445 });
446 }
447
448 pub fn set_g2_blocks(&self, blocks: Vec<BlockInfo>) {
450 self.tx.send_modify(|state| {
451 state.g2_blocks = blocks;
452 });
453 }
454
455 pub fn add_staged_blocks(
462 &self,
463 staged: Vec<BlockInfo>,
464 g3_remaining: usize,
465 layer_range: Option<std::ops::Range<usize>>,
466 ) {
467 self.tx.send_modify(|state| {
468 state.g2_blocks.extend(staged);
469 state.g3_pending = g3_remaining;
470 state.ready_layer_range = layer_range;
471 if g3_remaining == 0 && state.ready_layer_range.is_none() {
472 state.phase = SessionPhase::Ready;
474 }
475 });
476 }
477
478 pub fn set_failed(&self) {
480 self.tx.send_modify(|state| {
481 state.phase = SessionPhase::Failed;
482 });
483 }
484}
485
486pub fn session_handle_state_channel()
488-> (SessionHandleStateTx, watch::Receiver<SessionStateSnapshot>) {
489 let initial = SessionStateSnapshot {
490 phase: SessionPhase::Searching,
491 control_role: ControlRole::Controllee,
492 g2_blocks: Vec::new(),
493 g3_pending: 0,
494 ready_layer_range: None,
495 };
496 let (tx, rx) = watch::channel(initial);
497 (SessionHandleStateTx::new(tx), rx)
498}
499
500#[cfg(test)]
501mod tests {
502 use super::*;
503 use dashmap::DashMap;
504
505 fn create_test_transport() -> Arc<MessageTransport> {
506 Arc::new(MessageTransport::local(
507 Arc::new(DashMap::new()),
508 Arc::new(DashMap::new()),
509 ))
510 }
511
512 #[test]
513 fn test_session_handle_state_channel() {
514 let (tx, rx) = session_handle_state_channel();
515
516 let state = rx.borrow().clone();
518 assert_eq!(state.phase, SessionPhase::Searching);
519 assert_eq!(state.control_role, ControlRole::Controllee);
520 assert!(state.g2_blocks.is_empty());
521
522 tx.set_phase(SessionPhase::Ready);
524 let state = rx.borrow().clone();
525 assert_eq!(state.phase, SessionPhase::Ready);
526 }
527
528 #[test]
529 fn test_session_handle_creation() {
530 let (_, rx) = session_handle_state_channel();
531 let transport = create_test_transport();
532 let session_id = SessionId::new_v4();
533 let remote_id = InstanceId::new_v4();
534 let local_id = InstanceId::new_v4();
535
536 let handle = SessionHandle::new(session_id, remote_id, local_id, transport, rx);
537
538 assert_eq!(handle.session_id(), session_id);
539 assert_eq!(handle.remote_instance(), remote_id);
540 assert_eq!(handle.local_instance(), local_id);
541 assert_eq!(handle.phase(), SessionPhase::Searching);
542 assert!(!handle.is_ready());
543 assert!(!handle.is_complete());
544 assert!(!handle.has_remote_metadata());
545 }
546
547 #[tokio::test]
548 async fn test_wait_for_ready() {
549 let (tx, rx) = session_handle_state_channel();
550 let transport = create_test_transport();
551 let session_id = SessionId::new_v4();
552
553 let mut handle = SessionHandle::new(
554 session_id,
555 InstanceId::new_v4(),
556 InstanceId::new_v4(),
557 transport,
558 rx,
559 );
560
561 tokio::spawn(async move {
563 tokio::time::sleep(tokio::time::Duration::from_millis(10)).await;
564 tx.set_phase(SessionPhase::Ready);
565 });
566
567 let state = handle.wait_for_ready().await.unwrap();
568 assert_eq!(state.phase, SessionPhase::Ready);
569 }
570
571 #[test]
572 fn test_add_staged_blocks() {
573 let (tx, rx) = session_handle_state_channel();
574
575 tx.update(SessionStateSnapshot {
577 phase: SessionPhase::Staging,
578 control_role: ControlRole::Controllee,
579 g2_blocks: Vec::new(),
580 g3_pending: 5,
581 ready_layer_range: None,
582 });
583
584 let state = rx.borrow().clone();
585 assert_eq!(state.g3_pending, 5);
586 assert!(state.g2_blocks.is_empty());
587
588 let block = BlockInfo {
590 block_id: 42,
591 sequence_hash: crate::SequenceHash::new(1, None, 100),
592 layout_handle: kvbm_physical::manager::LayoutHandle::new(0, 1),
593 };
594 tx.add_staged_blocks(vec![block], 0, None);
595
596 let state = rx.borrow().clone();
597 assert_eq!(state.g2_blocks.len(), 1);
598 assert_eq!(state.g3_pending, 0);
599 assert_eq!(state.phase, SessionPhase::Ready);
601 }
602
603 #[test]
604 fn test_set_failed() {
605 let (tx, rx) = session_handle_state_channel();
606
607 assert_eq!(rx.borrow().phase, SessionPhase::Searching);
609
610 tx.set_failed();
611
612 assert_eq!(rx.borrow().phase, SessionPhase::Failed);
613 }
614
615 #[tokio::test]
616 async fn test_wait_for_complete() {
617 let (tx, rx) = session_handle_state_channel();
618 let transport = create_test_transport();
619 let session_id = SessionId::new_v4();
620
621 let mut handle = SessionHandle::new(
622 session_id,
623 InstanceId::new_v4(),
624 InstanceId::new_v4(),
625 transport,
626 rx,
627 );
628
629 tokio::spawn(async move {
630 tokio::time::sleep(tokio::time::Duration::from_millis(10)).await;
631 tx.set_phase(SessionPhase::Complete);
632 });
633
634 let state = handle.wait_for_complete().await.unwrap();
635 assert_eq!(state.phase, SessionPhase::Complete);
636 assert!(handle.is_complete());
637 }
638}