Skip to main content

kvbm_engine/leader/session/
server_session.rs

1// SPDX-FileCopyrightText: Copyright (c) 2025-2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
2// SPDX-License-Identifier: Apache-2.0
3
4//! ServerSession: Merged server-side session for both G2-only and G2+G3 staging modes.
5//!
6//! Unifies `EndpointSession` (G2-only, Direct layout handles) and
7//! `ControllableSession` (G2+G3 staging, RoundRobin layout handles) into a
8//! single type that uses `SessionEndpoint` for the state machine.
9//!
10//! # Modes
11//!
12//! - **G2-only**: Blocks are already in G2 with pre-assigned layout handles.
13//!   Created via `ServerSession::new_g2_only()` with `Direct` metadata.
14//!   `TriggerStaging` is a no-op.
15//!
16//! - **Staging**: G3 blocks need to be staged to G2. Layout handles are
17//!   assigned round-robin across workers. Created with `RoundRobin` metadata
18//!   and optional `auto_stage`.
19//!
20//! # Lifecycle
21//!
22//! 1. Created with G2 blocks (and optionally G3 blocks)
23//! 2. If `auto_stage=true`, immediately stages G3→G2
24//! 3. Waits for peer to `Attach`
25//! 4. Sends `StateResponse` with block info
26//! 5. Responds to `TriggerStaging`, `BlocksPulled`, `Detach`, etc.
27//! 6. Completes when all blocks pulled or session closed
28
29use std::collections::HashMap;
30use std::ops::Range;
31use std::sync::Arc;
32
33use anyhow::Result;
34use tokio::sync::mpsc;
35use tracing::{debug, warn};
36
37use kvbm_physical::manager::LayoutHandle;
38
39use super::SessionId;
40use super::blocks::BlockHolder;
41use super::endpoint::SessionEndpoint;
42use super::messages::{BlockInfo, SessionMessage, SessionStateSnapshot};
43use super::staging;
44use super::state::{ControlRole, SessionPhase};
45use super::transport::MessageTransport;
46use crate::{G2, G3, InstanceId, SequenceHash, worker::group::ParallelWorkers};
47use kvbm_logical::manager::BlockManager;
48
49/// Block metadata strategy for mapping blocks to layout handles.
50///
51/// Unifies the two approaches from the former EndpointSession (Direct)
52/// and ControllableSession (RoundRobin).
53pub enum BlockMetadataMap {
54    /// Pre-assigned layout handles keyed by sequence hash.
55    /// Used for G2-only mode where the caller knows exactly which handle
56    /// each block should use.
57    Direct(HashMap<SequenceHash, LayoutHandle>),
58
59    /// Worker layout handles for round-robin assignment.
60    /// Used for staging mode where blocks are distributed across workers.
61    RoundRobin(Vec<LayoutHandle>),
62}
63
64impl BlockMetadataMap {
65    /// Build `BlockInfo` list from the current G2 blocks.
66    fn build_block_infos(&self, g2_blocks: &BlockHolder<G2>) -> Vec<BlockInfo> {
67        match self {
68            BlockMetadataMap::Direct(map) => g2_blocks
69                .blocks()
70                .iter()
71                .filter_map(|block| {
72                    let hash = block.sequence_hash();
73                    map.get(&hash).map(|&layout_handle| BlockInfo {
74                        block_id: block.block_id(),
75                        sequence_hash: hash,
76                        layout_handle,
77                    })
78                })
79                .collect(),
80
81            BlockMetadataMap::RoundRobin(handles) => {
82                if handles.is_empty() {
83                    return g2_blocks
84                        .blocks()
85                        .iter()
86                        .map(|b| BlockInfo {
87                            block_id: b.block_id(),
88                            sequence_hash: b.sequence_hash(),
89                            layout_handle: LayoutHandle::new(0, 0),
90                        })
91                        .collect();
92                }
93                g2_blocks
94                    .blocks()
95                    .iter()
96                    .enumerate()
97                    .map(|(i, b)| BlockInfo {
98                        block_id: b.block_id(),
99                        sequence_hash: b.sequence_hash(),
100                        layout_handle: handles[i % handles.len()],
101                    })
102                    .collect()
103            }
104        }
105    }
106
107    /// Assign a layout handle for a newly staged block at the given index.
108    fn assign_handle(&self, index: usize) -> LayoutHandle {
109        match self {
110            BlockMetadataMap::Direct(_) => {
111                // Direct mode shouldn't be staging, but provide a fallback
112                LayoutHandle::new(0, 0)
113            }
114            BlockMetadataMap::RoundRobin(handles) => {
115                if handles.is_empty() {
116                    LayoutHandle::new(0, 0)
117                } else {
118                    handles[index % handles.len()]
119                }
120            }
121        }
122    }
123
124    /// Remove entries for the given sequence hashes (Direct mode only).
125    fn remove_all(&mut self, hashes: &[SequenceHash]) {
126        if let BlockMetadataMap::Direct(map) = self {
127            for hash in hashes {
128                map.remove(hash);
129            }
130        }
131    }
132}
133
134/// Options for server session creation.
135#[derive(Debug, Clone)]
136pub struct ServerSessionOptions {
137    /// If true (default), immediately start G3→G2 staging.
138    /// If false, wait for controller to call trigger_staging().
139    pub auto_stage: bool,
140}
141
142impl Default for ServerSessionOptions {
143    fn default() -> Self {
144        Self { auto_stage: true }
145    }
146}
147
148/// Server-side session that holds blocks and exposes them for remote RDMA pull.
149///
150/// Merges the functionality of the former `EndpointSession` and `ControllableSession`.
151pub struct ServerSession {
152    /// State machine for the session protocol.
153    endpoint: SessionEndpoint,
154
155    /// G2 blocks held for RDMA pull (RAII - released on drop).
156    g2_blocks: BlockHolder<G2>,
157
158    /// Block metadata mapping (Direct or RoundRobin).
159    block_metadata: BlockMetadataMap,
160
161    /// G3 blocks pending staging (empty in G2-only mode).
162    g3_blocks: BlockHolder<G3>,
163
164    /// G2 manager for staging (only needed when G3 blocks present).
165    g2_manager: Option<Arc<BlockManager<G2>>>,
166
167    /// Parallel worker for G3→G2 staging.
168    parallel_worker: Option<Arc<dyn ParallelWorkers>>,
169
170    /// Channel for receiving local commands.
171    cmd_rx: mpsc::Receiver<ServerSessionCommand>,
172
173    /// Session options.
174    options: ServerSessionOptions,
175
176    /// Staging state tracking.
177    staging_started: bool,
178    staging_complete: bool,
179}
180
181/// Handle for local caller to control a ServerSession.
182///
183/// Used to send layer notifications or close the session.
184/// When dropped, the session continues until peer detaches or channel closes.
185#[derive(Clone)]
186pub struct ServerSessionHandle {
187    session_id: SessionId,
188    local_instance: InstanceId,
189    cmd_tx: mpsc::Sender<ServerSessionCommand>,
190}
191
192/// Commands that can be sent to a ServerSession via its handle.
193#[derive(Debug)]
194pub enum ServerSessionCommand {
195    /// Notify that specific layers are ready for transfer.
196    NotifyLayersReady { layer_range: Range<usize> },
197    /// Close the session gracefully.
198    Close,
199}
200
201impl ServerSession {
202    /// Create a new ServerSession for G2-only mode.
203    ///
204    /// Blocks are already in G2 with pre-assigned layout handles.
205    pub fn new_g2_only(
206        endpoint: SessionEndpoint,
207        g2_blocks: BlockHolder<G2>,
208        block_metadata: HashMap<SequenceHash, LayoutHandle>,
209        cmd_rx: mpsc::Receiver<ServerSessionCommand>,
210    ) -> Self {
211        Self {
212            endpoint,
213            g2_blocks,
214            block_metadata: BlockMetadataMap::Direct(block_metadata),
215            g3_blocks: BlockHolder::empty(),
216            g2_manager: None,
217            parallel_worker: None,
218            cmd_rx,
219            options: ServerSessionOptions { auto_stage: false },
220            staging_started: false,
221            staging_complete: false,
222        }
223    }
224
225    /// Create a new ServerSession with G3→G2 staging capability.
226    #[allow(clippy::too_many_arguments)]
227    pub fn new_with_staging(
228        endpoint: SessionEndpoint,
229        g2_blocks: BlockHolder<G2>,
230        g3_blocks: BlockHolder<G3>,
231        worker_handles: Vec<LayoutHandle>,
232        g2_manager: Arc<BlockManager<G2>>,
233        parallel_worker: Option<Arc<dyn ParallelWorkers>>,
234        cmd_rx: mpsc::Receiver<ServerSessionCommand>,
235        options: ServerSessionOptions,
236    ) -> Self {
237        Self {
238            endpoint,
239            g2_blocks,
240            block_metadata: BlockMetadataMap::RoundRobin(worker_handles),
241            g3_blocks,
242            g2_manager: Some(g2_manager),
243            parallel_worker,
244            cmd_rx,
245            options,
246            staging_started: false,
247            staging_complete: false,
248        }
249    }
250
251    /// Run the session message loop.
252    pub async fn run(mut self) -> Result<()> {
253        debug!(
254            session_id = %self.endpoint.session_id(),
255            g2 = self.g2_blocks.count(),
256            g3 = self.g3_blocks.count(),
257            "ServerSession starting"
258        );
259
260        // Set initial phase
261        if self.g2_blocks.count() > 0 || self.g3_blocks.count() > 0 {
262            self.endpoint.set_phase(SessionPhase::Holding);
263        }
264
265        // Auto-stage if enabled and we have G3 blocks
266        if self.options.auto_stage && !self.g3_blocks.is_empty() && self.parallel_worker.is_some() {
267            self.endpoint.set_phase(SessionPhase::Staging);
268            self.staging_started = true;
269            self.execute_staging().await?;
270        }
271
272        self.update_phase();
273
274        loop {
275            tokio::select! {
276                msg = self.endpoint.recv() => {
277                    match msg {
278                        Some(msg) => {
279                            if !self.handle_message(msg).await? {
280                                break;
281                            }
282                        }
283                        None => {
284                            debug!(
285                                session_id = %self.endpoint.session_id(),
286                                "Message channel closed"
287                            );
288                            break;
289                        }
290                    }
291                }
292
293                cmd = self.cmd_rx.recv() => {
294                    match cmd {
295                        Some(cmd) => {
296                            if !self.handle_command(cmd).await? {
297                                break;
298                            }
299                        }
300                        None => {
301                            debug!(
302                                session_id = %self.endpoint.session_id(),
303                                "Command channel closed"
304                            );
305                        }
306                    }
307                }
308            }
309        }
310
311        debug!(
312            session_id = %self.endpoint.session_id(),
313            phase = ?self.endpoint.phase(),
314            "ServerSession completed"
315        );
316
317        Ok(())
318    }
319
320    /// Handle an incoming SessionMessage.
321    ///
322    /// Returns `true` to continue, `false` to exit the loop.
323    async fn handle_message(&mut self, msg: SessionMessage) -> Result<bool> {
324        match msg {
325            SessionMessage::Attach { peer, as_role, .. } => {
326                debug!(
327                    session_id = %self.endpoint.session_id(),
328                    peer = %peer,
329                    role = ?as_role,
330                    "Peer attached"
331                );
332
333                self.endpoint.accept_attachment(peer, as_role.opposite());
334
335                // Update phase for attach
336                if self.endpoint.phase() == SessionPhase::Searching
337                    || self.endpoint.phase() == SessionPhase::Holding
338                {
339                    self.update_phase();
340                }
341
342                // Send current state
343                self.send_state_response(None).await?;
344            }
345
346            SessionMessage::TriggerStaging { .. } => {
347                self.handle_trigger_staging().await?;
348            }
349
350            SessionMessage::BlocksPulled { pulled_hashes, .. } => {
351                debug!(
352                    session_id = %self.endpoint.session_id(),
353                    count = pulled_hashes.len(),
354                    "Blocks pulled"
355                );
356
357                self.block_metadata.remove_all(&pulled_hashes);
358                self.g2_blocks.release(&pulled_hashes);
359
360                if self.g2_blocks.is_empty() && self.g3_blocks.is_empty() {
361                    self.endpoint.set_phase(SessionPhase::Complete);
362                    return Ok(false);
363                }
364            }
365
366            SessionMessage::YieldControl { peer, .. } => {
367                debug!(
368                    session_id = %self.endpoint.session_id(),
369                    peer = %peer,
370                    "Peer yielded control"
371                );
372                self.endpoint.set_control_role(ControlRole::Neutral);
373            }
374
375            SessionMessage::AcquireControl { peer, .. } => {
376                debug!(
377                    session_id = %self.endpoint.session_id(),
378                    peer = %peer,
379                    "Peer acquiring control"
380                );
381                self.endpoint.set_control_role(ControlRole::Controllee);
382            }
383
384            SessionMessage::Detach { peer, .. } => {
385                debug!(
386                    session_id = %self.endpoint.session_id(),
387                    peer = %peer,
388                    "Peer detached"
389                );
390                self.endpoint.detach();
391                self.endpoint.set_phase(SessionPhase::Complete);
392                return Ok(false);
393            }
394
395            SessionMessage::Close { .. } => {
396                debug!(
397                    session_id = %self.endpoint.session_id(),
398                    "Session closed"
399                );
400                self.endpoint.set_phase(SessionPhase::Complete);
401                return Ok(false);
402            }
403
404            SessionMessage::Error { message, .. } => {
405                warn!(
406                    session_id = %self.endpoint.session_id(),
407                    error = %message,
408                    "Received error"
409                );
410                self.endpoint.set_phase(SessionPhase::Failed);
411                return Ok(false);
412            }
413
414            // Ignore outbound-only messages
415            SessionMessage::StateResponse { .. }
416            | SessionMessage::BlocksStaged { .. }
417            | SessionMessage::HoldBlocks { .. }
418            | SessionMessage::ReleaseBlocks { .. } => {}
419        }
420
421        Ok(true)
422    }
423
424    /// Handle a local command.
425    ///
426    /// Returns `true` to continue, `false` to exit the loop.
427    async fn handle_command(&mut self, cmd: ServerSessionCommand) -> Result<bool> {
428        match cmd {
429            ServerSessionCommand::NotifyLayersReady { layer_range } => {
430                debug!(
431                    session_id = %self.endpoint.session_id(),
432                    layer_range = ?layer_range,
433                    "Notifying layers ready"
434                );
435                self.send_blocks_staged(Some(layer_range)).await?;
436            }
437            ServerSessionCommand::Close => {
438                debug!(
439                    session_id = %self.endpoint.session_id(),
440                    "Local close requested"
441                );
442                self.endpoint.set_phase(SessionPhase::Complete);
443
444                if self.endpoint.is_attached() {
445                    let msg = SessionMessage::Close {
446                        session_id: self.endpoint.session_id(),
447                    };
448                    self.endpoint.send(msg).await?;
449                }
450                return Ok(false);
451            }
452        }
453
454        Ok(true)
455    }
456
457    /// Handle trigger staging request (idempotent).
458    async fn handle_trigger_staging(&mut self) -> Result<()> {
459        if self.staging_started {
460            return Ok(());
461        }
462
463        if self.g3_blocks.is_empty() {
464            // No-op for G2-only mode
465            debug!(
466                session_id = %self.endpoint.session_id(),
467                "TriggerStaging ignored (no G3 blocks)"
468            );
469            return Ok(());
470        }
471
472        if self.parallel_worker.is_none() {
473            if self.endpoint.is_attached() {
474                let error_msg = SessionMessage::Error {
475                    session_id: self.endpoint.session_id(),
476                    message: "No parallel worker available for G3->G2 staging".to_string(),
477                };
478                self.endpoint.send(error_msg).await?;
479            }
480            return Ok(());
481        }
482
483        self.endpoint.set_phase(SessionPhase::Staging);
484        self.staging_started = true;
485
486        let staged_info = self.execute_staging().await?;
487
488        self.update_phase();
489
490        // Notify peer of newly staged blocks (if attached)
491        if self.endpoint.is_attached() {
492            let msg = SessionMessage::BlocksStaged {
493                session_id: self.endpoint.session_id(),
494                staged_blocks: staged_info,
495                remaining: self.g3_blocks.count(),
496                layer_range: None,
497            };
498            self.endpoint.send(msg).await?;
499        }
500
501        Ok(())
502    }
503
504    /// Execute G3→G2 staging.
505    ///
506    /// Returns BlockInfo for newly staged blocks.
507    async fn execute_staging(&mut self) -> Result<Vec<BlockInfo>> {
508        let parallel_worker = self
509            .parallel_worker
510            .as_ref()
511            .ok_or_else(|| anyhow::anyhow!("ParallelWorkers required for G3→G2 staging"))?;
512
513        let g2_manager = self
514            .g2_manager
515            .as_ref()
516            .ok_or_else(|| anyhow::anyhow!("G2 manager required for staging"))?;
517
518        if self.g3_blocks.is_empty() {
519            self.staging_complete = true;
520            return Ok(Vec::new());
521        }
522
523        let result =
524            staging::stage_g3_to_g2(&self.g3_blocks, g2_manager, &**parallel_worker).await?;
525
526        // Build BlockInfo for newly staged blocks
527        let starting_index = self.g2_blocks.count();
528        let staged_info: Vec<BlockInfo> = result
529            .new_g2_blocks
530            .iter()
531            .enumerate()
532            .map(|(i, b)| BlockInfo {
533                block_id: b.block_id(),
534                sequence_hash: b.sequence_hash(),
535                layout_handle: self.block_metadata.assign_handle(starting_index + i),
536            })
537            .collect();
538
539        // Clear G3, extend G2
540        let _ = self.g3_blocks.take_all();
541        self.g2_blocks.extend(result.new_g2_blocks);
542
543        self.staging_complete = true;
544
545        Ok(staged_info)
546    }
547
548    /// Update phase based on current state.
549    fn update_phase(&mut self) {
550        if self.endpoint.phase() == SessionPhase::Complete
551            || self.endpoint.phase() == SessionPhase::Failed
552        {
553            return;
554        }
555
556        if self.g3_blocks.is_empty() && (self.staging_complete || !self.staging_started) {
557            self.endpoint.set_phase(SessionPhase::Ready);
558        } else if self.staging_started && !self.staging_complete {
559            self.endpoint.set_phase(SessionPhase::Staging);
560        }
561    }
562
563    /// Send a StateResponse to the attached peer.
564    async fn send_state_response(&self, layer_range: Option<Range<usize>>) -> Result<()> {
565        let state = self.build_state_snapshot(layer_range);
566        let msg = SessionMessage::StateResponse {
567            session_id: self.endpoint.session_id(),
568            state,
569        };
570        self.endpoint.send(msg).await
571    }
572
573    /// Send a BlocksStaged message with optional layer range.
574    async fn send_blocks_staged(&self, layer_range: Option<Range<usize>>) -> Result<()> {
575        let blocks = self.block_metadata.build_block_infos(&self.g2_blocks);
576        let msg = SessionMessage::BlocksStaged {
577            session_id: self.endpoint.session_id(),
578            staged_blocks: blocks,
579            remaining: 0,
580            layer_range,
581        };
582        self.endpoint.send(msg).await
583    }
584
585    /// Build a state snapshot.
586    fn build_state_snapshot(&self, layer_range: Option<Range<usize>>) -> SessionStateSnapshot {
587        SessionStateSnapshot {
588            phase: self.endpoint.phase(),
589            control_role: self.endpoint.control_role(),
590            g2_blocks: self.block_metadata.build_block_infos(&self.g2_blocks),
591            g3_pending: self.g3_blocks.count(),
592            ready_layer_range: layer_range,
593        }
594    }
595
596    /// Get the session ID.
597    pub fn session_id(&self) -> SessionId {
598        self.endpoint.session_id()
599    }
600}
601
602impl ServerSessionHandle {
603    /// Create a new server session handle.
604    pub fn new(
605        session_id: SessionId,
606        local_instance: InstanceId,
607        cmd_tx: mpsc::Sender<ServerSessionCommand>,
608    ) -> Self {
609        Self {
610            session_id,
611            local_instance,
612            cmd_tx,
613        }
614    }
615
616    /// Get the session ID.
617    pub fn session_id(&self) -> SessionId {
618        self.session_id
619    }
620
621    /// Get the local instance ID.
622    pub fn local_instance(&self) -> InstanceId {
623        self.local_instance
624    }
625
626    /// Notify attached controller that layers are ready.
627    pub async fn notify_layers_ready(&self, layer_range: Range<usize>) -> Result<()> {
628        self.cmd_tx
629            .send(ServerSessionCommand::NotifyLayersReady { layer_range })
630            .await
631            .map_err(|_| anyhow::anyhow!("Session command channel closed"))
632    }
633
634    /// Close the session gracefully.
635    pub async fn close(&self) -> Result<()> {
636        self.cmd_tx
637            .send(ServerSessionCommand::Close)
638            .await
639            .map_err(|_| anyhow::anyhow!("Session command channel closed"))
640    }
641}
642
643/// Create a ServerSession in G2-only mode with its handle.
644///
645/// This is the replacement for `create_endpoint_session`.
646pub fn create_server_session(
647    session_id: SessionId,
648    instance_id: InstanceId,
649    blocks: BlockHolder<G2>,
650    layout_handles: Vec<LayoutHandle>,
651    sequence_hashes: Vec<SequenceHash>,
652    transport: Arc<MessageTransport>,
653    msg_rx: mpsc::Receiver<SessionMessage>,
654) -> (ServerSession, ServerSessionHandle) {
655    let (cmd_tx, cmd_rx) = mpsc::channel(16);
656
657    let block_metadata: HashMap<SequenceHash, LayoutHandle> =
658        sequence_hashes.into_iter().zip(layout_handles).collect();
659
660    let endpoint = SessionEndpoint::new(session_id, instance_id, transport, msg_rx);
661
662    let session = ServerSession::new_g2_only(endpoint, blocks, block_metadata, cmd_rx);
663
664    let handle = ServerSessionHandle::new(session_id, instance_id, cmd_tx);
665
666    (session, handle)
667}
668
669#[cfg(test)]
670mod tests {
671    use super::*;
672    use crate::leader::session::SessionMessageTx;
673    use dashmap::DashMap;
674    use tokio::sync::mpsc;
675
676    fn create_test_transport() -> Arc<MessageTransport> {
677        Arc::new(MessageTransport::local(
678            Arc::new(DashMap::new()),
679            Arc::new(DashMap::new()),
680        ))
681    }
682
683    #[tokio::test]
684    async fn test_handle_creation() {
685        let (cmd_tx, _cmd_rx) = mpsc::channel(16);
686        let session_id = SessionId::new_v4();
687        let instance_id = InstanceId::new_v4();
688
689        let handle = ServerSessionHandle::new(session_id, instance_id, cmd_tx);
690
691        assert_eq!(handle.session_id(), session_id);
692        assert_eq!(handle.local_instance(), instance_id);
693    }
694
695    #[tokio::test]
696    async fn test_notify_layers_ready() {
697        let (cmd_tx, mut cmd_rx) = mpsc::channel(16);
698        let session_id = SessionId::new_v4();
699        let instance_id = InstanceId::new_v4();
700
701        let handle = ServerSessionHandle::new(session_id, instance_id, cmd_tx);
702
703        handle.notify_layers_ready(0..1).await.unwrap();
704
705        let cmd = cmd_rx.recv().await.unwrap();
706        match cmd {
707            ServerSessionCommand::NotifyLayersReady { layer_range } => {
708                assert_eq!(layer_range, 0..1);
709            }
710            _ => panic!("Unexpected command"),
711        }
712    }
713
714    #[tokio::test]
715    async fn test_handle_close() {
716        let (cmd_tx, mut cmd_rx) = mpsc::channel(16);
717        let session_id = SessionId::new_v4();
718        let instance_id = InstanceId::new_v4();
719
720        let handle = ServerSessionHandle::new(session_id, instance_id, cmd_tx);
721
722        handle.close().await.unwrap();
723
724        let cmd = cmd_rx.recv().await.unwrap();
725        assert!(matches!(cmd, ServerSessionCommand::Close));
726    }
727
728    #[tokio::test]
729    async fn test_create_server_session() {
730        let session_id = SessionId::new_v4();
731        let instance_id = InstanceId::new_v4();
732        let transport = create_test_transport();
733        let (_msg_tx, msg_rx) = mpsc::channel(16);
734
735        let blocks = BlockHolder::empty();
736
737        let (_session, handle) = create_server_session(
738            session_id,
739            instance_id,
740            blocks,
741            vec![],
742            vec![],
743            transport,
744            msg_rx,
745        );
746
747        assert_eq!(handle.session_id(), session_id);
748        assert_eq!(handle.local_instance(), instance_id);
749    }
750
751    #[tokio::test]
752    async fn test_attach_sends_state_response() {
753        let session_id = SessionId::new_v4();
754        let instance_id = InstanceId::new_v4();
755        let peer_id = InstanceId::new_v4();
756
757        // Create transport with a session channel to capture responses
758        let session_sessions: Arc<DashMap<SessionId, SessionMessageTx>> = Arc::new(DashMap::new());
759        let transport = Arc::new(MessageTransport::local(
760            Arc::new(DashMap::new()),
761            session_sessions.clone(),
762        ));
763
764        // Register a receiver for the peer's session (where StateResponse is sent)
765        let peer_session_id = SessionId::new_v4(); // peer's session
766        let (peer_tx, mut peer_rx) = mpsc::channel::<SessionMessage>(16);
767        session_sessions.insert(session_id, peer_tx);
768
769        let (msg_tx, msg_rx) = mpsc::channel(16);
770        let (_cmd_tx, cmd_rx) = mpsc::channel(16);
771
772        let endpoint = SessionEndpoint::new(session_id, instance_id, transport, msg_rx);
773        let session =
774            ServerSession::new_g2_only(endpoint, BlockHolder::empty(), HashMap::new(), cmd_rx);
775
776        // Spawn session
777        let session_task = tokio::spawn(session.run());
778
779        // Send attach message
780        msg_tx
781            .send(SessionMessage::Attach {
782                peer: peer_id,
783                session_id,
784                as_role: ControlRole::Controller,
785            })
786            .await
787            .unwrap();
788
789        // Read the StateResponse
790        let response = tokio::time::timeout(std::time::Duration::from_secs(1), peer_rx.recv())
791            .await
792            .expect("timeout")
793            .expect("channel closed");
794
795        match response {
796            SessionMessage::StateResponse { state, .. } => {
797                assert_eq!(state.phase, SessionPhase::Ready);
798                assert_eq!(state.control_role, ControlRole::Controllee);
799            }
800            other => panic!("Expected StateResponse, got {:?}", other),
801        }
802
803        // Close session
804        msg_tx
805            .send(SessionMessage::Close { session_id })
806            .await
807            .unwrap();
808
809        let _ = tokio::time::timeout(std::time::Duration::from_secs(1), session_task).await;
810
811        let _ = peer_session_id;
812    }
813
814    #[tokio::test]
815    async fn test_g2_only_ready_on_attach() {
816        let session_id = SessionId::new_v4();
817        let instance_id = InstanceId::new_v4();
818        let peer_id = InstanceId::new_v4();
819
820        let session_sessions: Arc<DashMap<SessionId, SessionMessageTx>> = Arc::new(DashMap::new());
821        let transport = Arc::new(MessageTransport::local(
822            Arc::new(DashMap::new()),
823            session_sessions.clone(),
824        ));
825
826        let (peer_tx, mut peer_rx) = mpsc::channel::<SessionMessage>(16);
827        session_sessions.insert(session_id, peer_tx);
828
829        let (msg_tx, msg_rx) = mpsc::channel(16);
830        let (_cmd_tx, cmd_rx) = mpsc::channel(16);
831
832        let endpoint = SessionEndpoint::new(session_id, instance_id, transport, msg_rx);
833        // G2-only mode, no G3 blocks
834        let session =
835            ServerSession::new_g2_only(endpoint, BlockHolder::empty(), HashMap::new(), cmd_rx);
836
837        let session_task = tokio::spawn(session.run());
838
839        msg_tx
840            .send(SessionMessage::Attach {
841                peer: peer_id,
842                session_id,
843                as_role: ControlRole::Controller,
844            })
845            .await
846            .unwrap();
847
848        let response = tokio::time::timeout(std::time::Duration::from_secs(1), peer_rx.recv())
849            .await
850            .expect("timeout")
851            .expect("channel closed");
852
853        // G2-only with no blocks → Ready phase immediately
854        match response {
855            SessionMessage::StateResponse { state, .. } => {
856                assert_eq!(state.phase, SessionPhase::Ready);
857                assert_eq!(state.g3_pending, 0);
858            }
859            other => panic!("Expected StateResponse, got {:?}", other),
860        }
861
862        msg_tx
863            .send(SessionMessage::Close { session_id })
864            .await
865            .unwrap();
866        let _ = tokio::time::timeout(std::time::Duration::from_secs(1), session_task).await;
867    }
868
869    #[tokio::test]
870    async fn test_trigger_staging_no_g3_noop() {
871        let session_id = SessionId::new_v4();
872        let instance_id = InstanceId::new_v4();
873        let peer_id = InstanceId::new_v4();
874
875        let session_sessions: Arc<DashMap<SessionId, SessionMessageTx>> = Arc::new(DashMap::new());
876        let transport = Arc::new(MessageTransport::local(
877            Arc::new(DashMap::new()),
878            session_sessions.clone(),
879        ));
880
881        let (peer_tx, mut peer_rx) = mpsc::channel::<SessionMessage>(16);
882        session_sessions.insert(session_id, peer_tx);
883
884        let (msg_tx, msg_rx) = mpsc::channel(16);
885        let (_cmd_tx, cmd_rx) = mpsc::channel(16);
886
887        let endpoint = SessionEndpoint::new(session_id, instance_id, transport, msg_rx);
888        let session =
889            ServerSession::new_g2_only(endpoint, BlockHolder::empty(), HashMap::new(), cmd_rx);
890
891        let session_task = tokio::spawn(session.run());
892
893        // Attach
894        msg_tx
895            .send(SessionMessage::Attach {
896                peer: peer_id,
897                session_id,
898                as_role: ControlRole::Controller,
899            })
900            .await
901            .unwrap();
902
903        // Consume StateResponse
904        let _ = tokio::time::timeout(std::time::Duration::from_secs(1), peer_rx.recv())
905            .await
906            .expect("timeout");
907
908        // Send TriggerStaging - should be no-op (no G3 blocks)
909        msg_tx
910            .send(SessionMessage::TriggerStaging { session_id })
911            .await
912            .unwrap();
913
914        // Close and check no extra messages were sent
915        msg_tx
916            .send(SessionMessage::Close { session_id })
917            .await
918            .unwrap();
919
920        let _ = tokio::time::timeout(std::time::Duration::from_secs(1), session_task).await;
921    }
922
923    #[tokio::test]
924    async fn test_detach_completes_session() {
925        let session_id = SessionId::new_v4();
926        let instance_id = InstanceId::new_v4();
927        let peer_id = InstanceId::new_v4();
928
929        let session_sessions: Arc<DashMap<SessionId, SessionMessageTx>> = Arc::new(DashMap::new());
930        let transport = Arc::new(MessageTransport::local(
931            Arc::new(DashMap::new()),
932            session_sessions.clone(),
933        ));
934
935        let (peer_tx, mut _peer_rx) = mpsc::channel::<SessionMessage>(16);
936        session_sessions.insert(session_id, peer_tx);
937
938        let (msg_tx, msg_rx) = mpsc::channel(16);
939        let (_cmd_tx, cmd_rx) = mpsc::channel(16);
940
941        let endpoint = SessionEndpoint::new(session_id, instance_id, transport, msg_rx);
942        let session =
943            ServerSession::new_g2_only(endpoint, BlockHolder::empty(), HashMap::new(), cmd_rx);
944
945        let session_task = tokio::spawn(session.run());
946
947        // Attach then detach
948        msg_tx
949            .send(SessionMessage::Attach {
950                peer: peer_id,
951                session_id,
952                as_role: ControlRole::Controller,
953            })
954            .await
955            .unwrap();
956
957        msg_tx
958            .send(SessionMessage::Detach {
959                peer: peer_id,
960                session_id,
961            })
962            .await
963            .unwrap();
964
965        // Session should complete
966        let result = tokio::time::timeout(std::time::Duration::from_secs(1), session_task)
967            .await
968            .expect("timeout")
969            .expect("task panicked");
970
971        assert!(result.is_ok());
972    }
973
974    #[test]
975    fn test_block_metadata_direct_build_infos() {
976        let hash1 = SequenceHash::new(1, None, 100);
977        let hash2 = SequenceHash::new(2, None, 200);
978
979        let mut map = HashMap::new();
980        map.insert(hash1, LayoutHandle::new(0, 1));
981        map.insert(hash2, LayoutHandle::new(0, 2));
982
983        let metadata = BlockMetadataMap::Direct(map);
984
985        // Empty holder
986        let holder = BlockHolder::<G2>::empty();
987        let infos = metadata.build_block_infos(&holder);
988        assert!(infos.is_empty());
989    }
990
991    #[test]
992    fn test_block_metadata_round_robin_empty_handles() {
993        let metadata = BlockMetadataMap::RoundRobin(vec![]);
994        let holder = BlockHolder::<G2>::empty();
995        let infos = metadata.build_block_infos(&holder);
996        assert!(infos.is_empty());
997    }
998
999    #[test]
1000    fn test_block_metadata_assign_handle() {
1001        let h0 = LayoutHandle::new(0, 10);
1002        let h1 = LayoutHandle::new(1, 20);
1003        let metadata = BlockMetadataMap::RoundRobin(vec![h0, h1]);
1004
1005        assert_eq!(metadata.assign_handle(0), h0);
1006        assert_eq!(metadata.assign_handle(1), h1);
1007        assert_eq!(metadata.assign_handle(2), h0); // wraps around
1008    }
1009
1010    #[test]
1011    fn test_block_metadata_remove_all() {
1012        let hash1 = SequenceHash::new(1, None, 100);
1013        let hash2 = SequenceHash::new(2, None, 200);
1014
1015        let mut map = HashMap::new();
1016        map.insert(hash1, LayoutHandle::new(0, 1));
1017        map.insert(hash2, LayoutHandle::new(0, 2));
1018
1019        let mut metadata = BlockMetadataMap::Direct(map);
1020        metadata.remove_all(&[hash1]);
1021
1022        // Verify hash1 was removed
1023        if let BlockMetadataMap::Direct(ref inner) = metadata {
1024            assert!(!inner.contains_key(&hash1));
1025            assert!(inner.contains_key(&hash2));
1026        }
1027    }
1028}