Skip to main content

kvbm_engine/leader/session/
handle.rs

1// SPDX-FileCopyrightText: Copyright (c) 2025-2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
2// SPDX-License-Identifier: Apache-2.0
3
4//! SessionHandle: Unified handle for controlling a remote session.
5//!
6//! This is the unified replacement for `RemoteSessionHandle` that uses the
7//! new session model types (`SessionPhase`, `ControlRole`, `SessionStateSnapshot`).
8//!
9//! Key improvements over RemoteSessionHandle:
10//! - Uses unified `SessionPhase` and `ControlRole` enums
11//! - Supports bidirectional control transfer (yield/acquire)
12//! - Uses `SessionStateSnapshot` for state observation
13//! - Same RDMA support via `ParallelWorker`
14
15use 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
29/// Handle for controlling a remote session.
30///
31/// Created by attaching to a remote session. Provides methods to:
32/// - Query and observe session state
33/// - Issue control commands (trigger staging, release blocks)
34/// - Transfer control bidirectionally (yield/acquire)
35/// - Pull blocks via RDMA
36///
37/// ## Usage
38///
39/// ```ignore
40/// // Attach to remote session
41/// let mut handle = leader.attach_session(remote_id, session_id).await?;
42///
43/// // Wait for initial state
44/// let state = handle.wait_for_ready().await?;
45///
46/// // Trigger staging if needed
47/// if state.g3_pending > 0 {
48///     handle.trigger_staging().await?;
49///     handle.wait_for_ready().await?;
50/// }
51///
52/// // Pull blocks via RDMA
53/// let notification = handle.pull_blocks_rdma(&state.g2_blocks, &local_block_ids).await?;
54/// notification.await?;
55///
56/// // Notify remote and detach
57/// handle.mark_blocks_pulled(hashes).await?;
58/// handle.detach().await?;
59/// ```
60pub struct SessionHandle {
61    session_id: SessionId,
62    remote_instance: InstanceId,
63    local_instance: InstanceId,
64    transport: Arc<MessageTransport>,
65
66    // State observation
67    state_rx: watch::Receiver<SessionStateSnapshot>,
68
69    // RDMA transfer support
70    parallel_worker: Option<Arc<dyn ParallelWorkers>>,
71}
72
73impl SessionHandle {
74    /// Create a new session handle.
75    ///
76    /// Note: Currently unused during incremental migration. Will be used once
77    /// existing session implementations are fully migrated to the new model.
78    #[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    /// Add RDMA support to this handle.
97    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    // =========================================================================
103    // Identity
104    // =========================================================================
105
106    /// Get the session ID.
107    pub fn session_id(&self) -> SessionId {
108        self.session_id
109    }
110
111    /// Get the remote instance ID.
112    pub fn remote_instance(&self) -> InstanceId {
113        self.remote_instance
114    }
115
116    /// Get the local instance ID.
117    pub fn local_instance(&self) -> InstanceId {
118        self.local_instance
119    }
120
121    // =========================================================================
122    // State Observation
123    // =========================================================================
124
125    /// Get the current state snapshot (non-blocking).
126    pub fn current_state(&self) -> SessionStateSnapshot {
127        self.state_rx.borrow().clone()
128    }
129
130    /// Get the current phase.
131    pub fn phase(&self) -> SessionPhase {
132        self.state_rx.borrow().phase
133    }
134
135    /// Get the current control role of the remote session.
136    pub fn remote_control_role(&self) -> ControlRole {
137        self.state_rx.borrow().control_role
138    }
139
140    /// Check if state has changed since last read.
141    pub fn has_changed(&self) -> bool {
142        self.state_rx.has_changed().unwrap_or(false)
143    }
144
145    /// Wait for state to change.
146    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    /// Wait for the session to reach Ready phase (all blocks in G2).
155    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    /// Wait for the session to complete.
169    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    /// Check if the session is complete.
178    pub fn is_complete(&self) -> bool {
179        self.state_rx.borrow().phase.is_terminal()
180    }
181
182    /// Check if the session is ready (all blocks in G2).
183    pub fn is_ready(&self) -> bool {
184        self.state_rx.borrow().phase == SessionPhase::Ready
185    }
186
187    /// Get G2 blocks from current state.
188    pub fn get_g2_blocks(&self) -> Vec<BlockInfo> {
189        self.state_rx.borrow().g2_blocks.clone()
190    }
191
192    /// Get count of G3 blocks pending staging.
193    pub fn g3_pending_count(&self) -> usize {
194        self.state_rx.borrow().g3_pending
195    }
196
197    /// Get the layer range that is ready for transfer.
198    ///
199    /// Returns `None` if all layers are ready or layerwise tracking is not active.
200    /// Returns `Some(range)` if only specific layers are ready.
201    pub fn ready_layer_range(&self) -> Option<std::ops::Range<usize>> {
202        self.state_rx.borrow().ready_layer_range.clone()
203    }
204
205    // =========================================================================
206    // Control Commands
207    // =========================================================================
208
209    /// Trigger G3→G2 staging on the remote session.
210    ///
211    /// Idempotent - no-op if already staging or staged.
212    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    /// Notify remote that blocks have been pulled.
220    ///
221    /// Call after successfully pulling blocks via RDMA.
222    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    /// Detach from the session.
231    ///
232    /// Consumes the handle. The remote session will release remaining blocks.
233    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    // =========================================================================
242    // Control Transfer (Bidirectional)
243    // =========================================================================
244
245    /// Yield control to the remote peer.
246    ///
247    /// After yielding, this handle transitions to Neutral and the remote
248    /// can acquire control if desired.
249    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    /// Attempt to acquire control from the remote peer.
258    ///
259    /// Valid when remote is in Neutral state.
260    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    // =========================================================================
269    // RDMA Transfer Methods
270    // =========================================================================
271
272    /// Check if remote metadata has been imported.
273    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    /// Ensure remote metadata is imported (lazy loading).
281    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    /// Pull blocks from remote G2 to local G2 via RDMA.
304    ///
305    /// This method:
306    /// 1. Ensures remote metadata is imported
307    /// 2. Executes SPMD-aware transfer (worker N pulls from remote worker N)
308    /// 3. Returns notification that completes when all transfers done
309    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    /// Pull blocks with explicit metadata pre-import.
319    ///
320    /// Caller must have already ensured metadata is imported.
321    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    /// Pull blocks from remote G2 to local G2 via RDMA with transfer options.
359    ///
360    /// This method allows specifying transfer options like layer range for
361    /// layerwise transfer. Use this when you only want to pull specific layers.
362    ///
363    /// # Example
364    /// ```ignore
365    /// // Pull only layer 0
366    /// let notification = handle.pull_blocks_rdma_with_options(
367    ///     &state.g2_blocks,
368    ///     &local_block_ids,
369    ///     TransferOptions::builder().layer_range(0..1).build(),
370    /// ).await?;
371    /// notification.await?;
372    /// ```
373    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    /// Pull blocks with options and explicit metadata pre-import.
384    ///
385    /// Caller must have already ensured metadata is imported.
386    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
425/// Sender for state updates to SessionHandle.
426pub struct SessionHandleStateTx {
427    tx: watch::Sender<SessionStateSnapshot>,
428}
429
430impl SessionHandleStateTx {
431    /// Create a new state sender.
432    pub fn new(tx: watch::Sender<SessionStateSnapshot>) -> Self {
433        Self { tx }
434    }
435
436    /// Update state from a full snapshot.
437    pub fn update(&self, state: SessionStateSnapshot) {
438        let _ = self.tx.send(state);
439    }
440
441    /// Update phase only.
442    pub fn set_phase(&self, phase: SessionPhase) {
443        self.tx.send_modify(|state| {
444            state.phase = phase;
445        });
446    }
447
448    /// Update G2 blocks.
449    pub fn set_g2_blocks(&self, blocks: Vec<BlockInfo>) {
450        self.tx.send_modify(|state| {
451            state.g2_blocks = blocks;
452        });
453    }
454
455    /// Add newly staged blocks.
456    ///
457    /// # Arguments
458    /// * `staged` - Blocks that have been staged
459    /// * `g3_remaining` - Count of G3 blocks still pending
460    /// * `layer_range` - Optional layer range that is ready for transfer
461    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                // All blocks staged and no layer tracking = fully ready
473                state.phase = SessionPhase::Ready;
474            }
475        });
476    }
477
478    /// Set error/failed state.
479    pub fn set_failed(&self) {
480        self.tx.send_modify(|state| {
481            state.phase = SessionPhase::Failed;
482        });
483    }
484}
485
486/// Create a new session handle state channel.
487pub 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        // Initial state
517        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        // Update state
523        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        // Spawn task to update state
562        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        // Set initial g3 pending
576        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        // Add staged blocks with remaining = 0
589        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        // No layer range + g3_remaining == 0 → Ready
600        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        // Initially Searching
608        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}