1use anyhow::Result;
5use tokio::sync::{Mutex, mpsc, watch};
6use tokio::task::JoinHandle;
7
8use std::collections::{HashMap, HashSet};
9use std::sync::Arc;
10
11use crate::{
12 BlockId, G2, G3, InstanceId, SequenceHash, object::ObjectBlockOps,
13 worker::group::ParallelWorkers,
14};
15use kvbm_common::LogicalLayoutHandle;
16use kvbm_logical::{blocks::ImmutableBlock, manager::BlockManager};
17use kvbm_physical::transfer::TransferOptions;
18
19use super::staging;
20
21use super::{
22 super::{OnboardingStatus, SessionControl, StagingMode},
23 BlockHolder, SessionId,
24 messages::OnboardMessage,
25 transport::MessageTransport,
26};
27
28fn validate_contiguous_positions(seq_hashes: &[SequenceHash]) -> Result<()> {
33 if seq_hashes.len() <= 1 {
34 return Ok(());
35 }
36
37 let mut positions: Vec<u64> = seq_hashes.iter().map(|h| h.position()).collect();
39 positions.sort();
40
41 for window in positions.windows(2) {
43 if window[1] != window[0] + 1 {
44 anyhow::bail!(
45 "Position gap detected in remote blocks: {} -> {} (expected {}). \
46 This indicates a block ordering bug.",
47 window[0],
48 window[1],
49 window[0] + 1
50 );
51 }
52 }
53
54 Ok(())
55}
56
57#[derive(Default)]
62struct G4SearchState {
63 won_hashes: HashSet<SequenceHash>,
65 pending_load: HashSet<SequenceHash>,
67 failed_hashes: HashMap<SequenceHash, String>,
69 allocated_blocks: HashMap<SequenceHash, BlockId>,
71}
72
73impl G4SearchState {
74 fn new() -> Self {
75 Self::default()
76 }
77
78 #[expect(dead_code)]
80 fn clear(&mut self) {
81 self.won_hashes.clear();
82 self.pending_load.clear();
83 self.failed_hashes.clear();
84 self.allocated_blocks.clear();
85 }
86}
87
88pub struct InitiatorSession {
95 session_id: SessionId,
96 instance_id: InstanceId,
97 mode: StagingMode,
98 g2_manager: Arc<BlockManager<G2>>,
99 g3_manager: Option<Arc<BlockManager<G3>>>,
100 parallel_worker: Option<Arc<dyn ParallelWorkers>>,
101 transport: Arc<MessageTransport>,
102 status_tx: watch::Sender<OnboardingStatus>,
103
104 local_g2_blocks: BlockHolder<G2>,
106 local_g3_blocks: BlockHolder<G3>,
107
108 remote_g2_blocks: HashMap<InstanceId, Vec<BlockId>>, remote_g2_hashes: HashMap<InstanceId, Vec<SequenceHash>>, remote_g3_blocks: HashMap<InstanceId, Vec<SequenceHash>>, all_g2_blocks: Arc<Mutex<Option<Vec<ImmutableBlock<G2>>>>>,
115
116 control_rx: mpsc::Receiver<SessionControl>,
118
119 object_client: Option<Arc<dyn ObjectBlockOps>>,
122 g4_state: G4SearchState,
124 g4_rx: Option<mpsc::Receiver<OnboardMessage>>,
126 #[allow(dead_code)]
128 g4_task_handle: Option<JoinHandle<()>>,
129}
130
131impl InitiatorSession {
132 #[allow(clippy::too_many_arguments)]
134 pub(crate) fn new(
135 session_id: SessionId,
136 instance_id: InstanceId,
137 mode: StagingMode,
138 g2_manager: Arc<BlockManager<G2>>,
139 g3_manager: Option<Arc<BlockManager<G3>>>,
140 parallel_worker: Option<Arc<dyn ParallelWorkers>>,
141 transport: Arc<MessageTransport>,
142 status_tx: watch::Sender<OnboardingStatus>,
143 all_g2_blocks: Arc<Mutex<Option<Vec<ImmutableBlock<G2>>>>>,
144 control_rx: mpsc::Receiver<SessionControl>,
145 object_client: Option<Arc<dyn ObjectBlockOps>>,
146 ) -> Self {
147 Self {
148 session_id,
149 instance_id,
150 mode,
151 g2_manager,
152 g3_manager,
153 parallel_worker,
154 transport,
155 status_tx,
156 local_g2_blocks: BlockHolder::empty(),
157 local_g3_blocks: BlockHolder::empty(),
158 remote_g2_blocks: HashMap::new(),
159 remote_g2_hashes: HashMap::new(),
160 remote_g3_blocks: HashMap::new(),
161 all_g2_blocks,
162 control_rx,
163 object_client,
164 g4_state: G4SearchState::new(),
165 g4_rx: None,
166 g4_task_handle: None,
167 }
168 }
169
170 pub async fn run(
172 mut self,
173 mut rx: mpsc::Receiver<OnboardMessage>,
174 remote_leaders: Vec<InstanceId>,
175 sequence_hashes: Vec<SequenceHash>,
176 ) -> Result<()> {
177 tracing::debug!(
178 session_id = %self.session_id,
179 mode = ?self.mode,
180 num_hashes = sequence_hashes.len(),
181 num_remotes = remote_leaders.len(),
182 "Starting initiator session"
183 );
184
185 self.search_phase(&mut rx, &remote_leaders, &sequence_hashes)
187 .await?;
188
189 self.apply_find_policy(&sequence_hashes).await?;
192
193 tracing::debug!(
194 session_id = %self.session_id,
195 "search_phase complete, entering mode handler"
196 );
197
198 match self.mode {
200 StagingMode::Hold => {
201 tracing::debug!(session_id = %self.session_id, "Calling hold_mode()");
202 self.hold_mode().await?;
203 self.await_commands(rx).await?;
205 }
206 StagingMode::Prepare => {
207 self.prepare_mode(&mut rx).await?;
208 self.await_commands(rx).await?;
210 }
211 StagingMode::Full => {
212 self.full_mode(&mut rx).await?;
213 }
215 }
216
217 Ok(())
218 }
219
220 async fn search_phase(
222 &mut self,
223 rx: &mut mpsc::Receiver<OnboardMessage>,
224 remote_leaders: &[InstanceId],
225 sequence_hashes: &[SequenceHash],
226 ) -> Result<()> {
227 self.local_g2_blocks = BlockHolder::new(self.g2_manager.match_blocks(sequence_hashes));
229
230 let mut matched_hashes: HashSet<SequenceHash> =
231 self.local_g2_blocks.sequence_hashes().into_iter().collect();
232
233 if let Some(ref g3_manager) = self.g3_manager {
235 let remaining: Vec<_> = sequence_hashes
236 .iter()
237 .filter(|h| !matched_hashes.contains(h))
238 .copied()
239 .collect();
240
241 if !remaining.is_empty() {
242 self.local_g3_blocks = BlockHolder::new(g3_manager.match_blocks(&remaining));
243 for hash in self.local_g3_blocks.sequence_hashes() {
244 matched_hashes.insert(hash);
245 }
246 }
247 }
248
249 let has_object_client = self.object_client.is_some();
252 if matched_hashes.len() == sequence_hashes.len()
253 || (remote_leaders.is_empty() && !has_object_client)
254 {
255 return Ok(());
256 }
257
258 let remaining_hashes: Vec<_> = sequence_hashes
260 .iter()
261 .filter(|h| !matched_hashes.contains(h))
262 .copied()
263 .collect();
264
265 if remaining_hashes.is_empty() {
266 return Ok(());
267 }
268
269 self.status_tx.send(OnboardingStatus::Searching).ok();
270
271 for remote in remote_leaders {
273 let msg = OnboardMessage::CreateSession {
274 requester: self.instance_id,
275 session_id: self.session_id,
276 sequence_hashes: remaining_hashes.clone(),
277 };
278 self.transport.send(*remote, msg).await?;
279 }
280
281 let g4_tx = if self.object_client.is_some() && self.parallel_worker.is_some() {
284 let (tx, rx) = mpsc::channel(16);
285 self.g4_rx = Some(rx);
286 let handle = self.spawn_g4_search(remaining_hashes.clone(), tx.clone());
288 self.g4_task_handle = Some(handle);
289 Some(tx)
290 } else {
291 None
292 };
293
294 self.process_search_responses(rx, remote_leaders, &mut matched_hashes, g4_tx)
296 .await?;
297
298 Ok(())
299 }
300
301 async fn process_search_responses(
306 &mut self,
307 rx: &mut mpsc::Receiver<OnboardMessage>,
308 remote_leaders: &[InstanceId],
309 matched_hashes: &mut HashSet<SequenceHash>,
310 g4_tx: Option<mpsc::Sender<OnboardMessage>>,
311 ) -> Result<()> {
312 let mut pending_g2_responses = remote_leaders.len();
313 let mut pending_g3_responses: HashSet<InstanceId> =
314 remote_leaders.iter().copied().collect();
315 let mut pending_search_complete: HashSet<InstanceId> =
316 remote_leaders.iter().copied().collect();
317 let mut pending_acknowledgments: HashSet<InstanceId> = HashSet::new();
318
319 let mut pending_g4_search = self.g4_rx.is_some();
321 let mut pending_g4_load = false;
322
323 let is_complete = |pending_g2: usize,
325 pending_g3: &HashSet<InstanceId>,
326 pending_ack: &HashSet<InstanceId>,
327 pending_search: &HashSet<InstanceId>,
328 pending_g4_s: bool,
329 pending_g4_l: bool| {
330 pending_g2 == 0
331 && pending_g3.is_empty()
332 && pending_ack.is_empty()
333 && pending_search.is_empty()
334 && !pending_g4_s
335 && !pending_g4_l
336 };
337
338 loop {
339 if is_complete(
341 pending_g2_responses,
342 &pending_g3_responses,
343 &pending_acknowledgments,
344 &pending_search_complete,
345 pending_g4_search,
346 pending_g4_load,
347 ) {
348 tracing::debug!(
349 session_id = %self.session_id,
350 "All responses received (including G4), exiting search_phase"
351 );
352 break;
353 }
354
355 tokio::select! {
356 g4_msg = async {
358 if let Some(ref mut g4_rx) = self.g4_rx {
359 g4_rx.recv().await
360 } else {
361 std::future::pending::<Option<OnboardMessage>>().await
362 }
363 } => {
364 let Some(msg) = g4_msg else {
365 pending_g4_search = false;
367 pending_g4_load = false;
368 continue;
369 };
370
371 tracing::debug!(
372 session_id = %self.session_id,
373 msg = msg.variant_name(),
374 "process_search_responses received G4"
375 );
376
377 match msg {
378 OnboardMessage::G4Results { found_hashes, .. } => {
379 pending_g4_search = false;
380
381 let won_hashes = self.process_g4_results(found_hashes, matched_hashes);
383
384 if !won_hashes.is_empty()
386 && let Some(ref tx) = g4_tx {
387 self.load_g4_blocks(won_hashes, tx.clone()).await?;
388 pending_g4_load = true;
389 }
390 }
391 OnboardMessage::G4LoadComplete { success, failures, blocks, .. } => {
392 self.handle_g4_load_complete(success, failures, blocks);
393 pending_g4_load = false;
394 }
395 _ => {}
396 }
397 }
398
399 remote_msg = rx.recv() => {
401 let Some(msg) = remote_msg else {
402 break;
404 };
405
406 tracing::debug!(
407 session_id = %self.session_id,
408 msg = msg.variant_name(),
409 "process_search_responses received"
410 );
411
412 match msg {
413 OnboardMessage::G2Results {
414 responder,
415 sequence_hashes,
416 block_ids,
417 ..
418 } => {
419 tracing::debug!(
420 session_id = %self.session_id,
421 responder = %responder,
422 num_hashes = sequence_hashes.len(),
423 "Processing G2Results"
424 );
425
426 let mut hold_hashes = Vec::new();
428 let mut drop_hashes = Vec::new();
429
430 for (seq_hash, block_id) in sequence_hashes.iter().zip(block_ids.iter()) {
431 if matched_hashes.insert(*seq_hash) {
432 hold_hashes.push(*seq_hash);
433 self.remote_g2_blocks
434 .entry(responder)
435 .or_default()
436 .push(*block_id);
437 self.remote_g2_hashes
439 .entry(responder)
440 .or_default()
441 .push(*seq_hash);
442 } else {
443 drop_hashes.push(*seq_hash);
444 }
445 }
446
447 self.transport
449 .send(
450 responder,
451 OnboardMessage::HoldBlocks {
452 requester: self.instance_id,
453 session_id: self.session_id,
454 hold_hashes,
455 drop_hashes,
456 },
457 )
458 .await?;
459
460 pending_acknowledgments.insert(responder);
461 pending_g2_responses -= 1;
462 }
463 OnboardMessage::G3Results {
464 responder,
465 sequence_hashes,
466 ..
467 } => {
468 for seq_hash in sequence_hashes {
470 if matched_hashes.insert(seq_hash) {
471 self.remote_g3_blocks
472 .entry(responder)
473 .or_default()
474 .push(seq_hash);
475 }
476 }
477
478 pending_g3_responses.remove(&responder);
479 }
480 OnboardMessage::SearchComplete { responder, .. } => {
481 pending_search_complete.remove(&responder);
482 pending_g3_responses.remove(&responder);
484
485 tracing::debug!(
486 session_id = %self.session_id,
487 responder = %responder,
488 g2_pending = pending_g2_responses,
489 g3_pending = pending_g3_responses.len(),
490 ack_pending = pending_acknowledgments.len(),
491 search_pending = pending_search_complete.len(),
492 g4_search = pending_g4_search,
493 g4_load = pending_g4_load,
494 "SearchComplete"
495 );
496 }
497 OnboardMessage::Acknowledged { responder, .. } => {
498 pending_acknowledgments.remove(&responder);
499 }
500 _ => {}
501 }
502 }
503 }
504 }
505
506 Ok(())
507 }
508
509 async fn apply_find_policy(&mut self, sequence_hashes: &[SequenceHash]) -> Result<()> {
518 let mut matched_hashes: HashSet<SequenceHash> = HashSet::new();
520
521 for hash in self.local_g2_blocks.sequence_hashes() {
523 matched_hashes.insert(hash);
524 }
525
526 for hash in self.local_g3_blocks.sequence_hashes() {
528 matched_hashes.insert(hash);
529 }
530
531 for hashes in self.remote_g2_hashes.values() {
533 for hash in hashes {
534 matched_hashes.insert(*hash);
535 }
536 }
537
538 for hashes in self.remote_g3_blocks.values() {
540 for hash in hashes {
541 matched_hashes.insert(*hash);
542 }
543 }
544
545 for hash in &self.g4_state.won_hashes {
547 matched_hashes.insert(*hash);
548 }
549
550 let mut keep_count = 0;
552 for hash in sequence_hashes {
553 if matched_hashes.contains(hash) {
554 keep_count += 1;
555 } else {
556 break;
558 }
559 }
560
561 if keep_count == sequence_hashes.len() || keep_count == matched_hashes.len() {
563 tracing::debug!(
564 session_id = %self.session_id,
565 matched = keep_count,
566 total = sequence_hashes.len(),
567 "apply_find_policy: no trimming needed"
568 );
569 return Ok(());
570 }
571
572 let keep_hashes: Vec<SequenceHash> = sequence_hashes[..keep_count].to_vec();
574 let keep_set: HashSet<&SequenceHash> = keep_hashes.iter().collect();
575
576 tracing::debug!(
577 session_id = %self.session_id,
578 from = matched_hashes.len(),
579 to = keep_count,
580 first_hole = keep_count,
581 "apply_find_policy: trimming blocks"
582 );
583
584 self.local_g2_blocks.retain(&keep_hashes);
586 self.local_g3_blocks.retain(&keep_hashes);
587
588 for (remote_instance, block_ids) in &mut self.remote_g2_blocks {
590 let hashes = self.remote_g2_hashes.get_mut(remote_instance);
591 if let Some(hashes) = hashes {
592 let mut release_indices = Vec::new();
594 for (i, hash) in hashes.iter().enumerate() {
595 if !keep_set.contains(hash) {
596 release_indices.push(i);
597 }
598 }
599
600 let release_hashes: Vec<SequenceHash> =
602 release_indices.iter().map(|&i| hashes[i]).collect();
603
604 for i in release_indices.into_iter().rev() {
606 hashes.remove(i);
607 block_ids.remove(i);
608 }
609
610 if !release_hashes.is_empty() {
612 tracing::debug!(
613 session_id = %self.session_id,
614 count = release_hashes.len(),
615 instance = %remote_instance,
616 "Releasing G2 blocks beyond first hole"
617 );
618 self.transport
619 .send(
620 *remote_instance,
621 OnboardMessage::ReleaseBlocks {
622 requester: self.instance_id,
623 session_id: self.session_id,
624 release_hashes,
625 },
626 )
627 .await?;
628 }
629 }
630 }
631
632 for (remote_instance, hashes) in &mut self.remote_g3_blocks {
634 let release_hashes: Vec<SequenceHash> = hashes
636 .iter()
637 .filter(|h| !keep_set.contains(h))
638 .copied()
639 .collect();
640
641 hashes.retain(|h| keep_set.contains(h));
643
644 if !release_hashes.is_empty() {
646 tracing::debug!(
647 session_id = %self.session_id,
648 count = release_hashes.len(),
649 instance = %remote_instance,
650 "Releasing G3 blocks beyond first hole"
651 );
652 self.transport
653 .send(
654 *remote_instance,
655 OnboardMessage::ReleaseBlocks {
656 requester: self.instance_id,
657 session_id: self.session_id,
658 release_hashes,
659 },
660 )
661 .await?;
662 }
663 }
664
665 let g4_release_hashes: Vec<SequenceHash> = self
667 .g4_state
668 .won_hashes
669 .iter()
670 .filter(|h| !keep_set.contains(h))
671 .copied()
672 .collect();
673
674 if !g4_release_hashes.is_empty() {
675 tracing::debug!(
676 session_id = %self.session_id,
677 count = g4_release_hashes.len(),
678 "Releasing G4 blocks beyond first hole"
679 );
680
681 for hash in &g4_release_hashes {
682 self.g4_state.won_hashes.remove(hash);
684 self.g4_state.pending_load.remove(hash);
686 self.g4_state.allocated_blocks.remove(hash);
688 }
689 }
690
691 Ok(())
692 }
693
694 async fn hold_mode(&mut self) -> Result<()> {
696 let local_g2 = self.local_g2_blocks.count();
697 let local_g3 = self.local_g3_blocks.count();
698 let remote_g2: usize = self.remote_g2_blocks.values().map(|v| v.len()).sum();
699 let remote_g3: usize = self.remote_g3_blocks.values().map(|v| v.len()).sum();
700
701 let pending_g4 = self.g4_state.pending_load.len();
703 let loaded_g4 = self.g4_state.won_hashes.len();
704 let failed_g4 = self.g4_state.failed_hashes.len();
705
706 tracing::debug!(
707 session_id = %self.session_id,
708 local_g2,
709 local_g3,
710 remote_g2,
711 remote_g3,
712 pending_g4,
713 loaded_g4,
714 failed_g4,
715 "hold_mode"
716 );
717
718 self.status_tx
719 .send(OnboardingStatus::Holding {
720 local_g2,
721 local_g3,
722 remote_g2,
723 remote_g3,
724 pending_g4,
725 loaded_g4,
726 failed_g4,
727 })
728 .ok();
729
730 tracing::debug!(session_id = %self.session_id, "Sent Holding status");
731
732 Ok(())
733 }
734
735 async fn send_stage_and_wait_for_ready(
740 &mut self,
741 rx: &mut mpsc::Receiver<OnboardMessage>,
742 ) -> Result<()> {
743 if self.remote_g3_blocks.is_empty() {
744 return Ok(());
745 }
746
747 let remotes_with_g3: Vec<(InstanceId, Vec<SequenceHash>)> = self
749 .remote_g3_blocks
750 .iter()
751 .map(|(k, v)| (*k, v.clone()))
752 .collect();
753
754 for (remote, stage_hashes) in &remotes_with_g3 {
755 self.transport
756 .send(
757 *remote,
758 OnboardMessage::StageBlocks {
759 requester: self.instance_id,
760 session_id: self.session_id,
761 stage_hashes: stage_hashes.clone(),
762 },
763 )
764 .await?;
765 }
766
767 let mut pending: HashSet<InstanceId> = remotes_with_g3.iter().map(|(k, _)| *k).collect();
769
770 while !pending.is_empty() {
771 match rx.recv().await {
772 Some(OnboardMessage::BlocksReady {
773 responder,
774 sequence_hashes,
775 block_ids,
776 ..
777 }) => {
778 tracing::debug!(
779 session_id = %self.session_id,
780 responder = %responder,
781 count = block_ids.len(),
782 "Received BlocksReady"
783 );
784 self.remote_g2_blocks
785 .entry(responder)
786 .or_default()
787 .extend(block_ids);
788 self.remote_g2_hashes
789 .entry(responder)
790 .or_default()
791 .extend(sequence_hashes);
792 pending.remove(&responder);
793 }
794 Some(other) => {
795 tracing::warn!(
796 session_id = %self.session_id,
797 msg = other.variant_name(),
798 "Unexpected message while waiting for BlocksReady"
799 );
800 }
801 None => {
802 tracing::warn!(
803 session_id = %self.session_id,
804 "Channel closed while waiting for BlocksReady"
805 );
806 break;
807 }
808 }
809 }
810
811 Ok(())
812 }
813
814 async fn prepare_mode(&mut self, rx: &mut mpsc::Receiver<OnboardMessage>) -> Result<()> {
816 self.stage_local_g3_to_g2().await?;
818
819 self.send_stage_and_wait_for_ready(rx).await?;
821
822 let local_g2 = self.local_g2_blocks.count();
823 let remote_g2: usize = self.remote_g2_blocks.values().map(|v| v.len()).sum();
824
825 self.status_tx
826 .send(OnboardingStatus::Prepared {
827 local_g2,
828 remote_g2,
829 })
830 .ok();
831
832 Ok(())
833 }
834
835 async fn full_mode(&mut self, rx: &mut mpsc::Receiver<OnboardMessage>) -> Result<()> {
837 self.stage_local_g3_to_g2().await?;
839
840 self.send_stage_and_wait_for_ready(rx).await?;
842
843 self.pull_remote_blocks().await?;
845
846 self.consolidate_blocks().await;
848
849 let all_remotes: HashSet<InstanceId> = self
851 .remote_g2_blocks
852 .keys()
853 .chain(self.remote_g3_blocks.keys())
854 .copied()
855 .collect();
856
857 for remote in all_remotes {
858 self.transport
859 .send(
860 remote,
861 OnboardMessage::CloseSession {
862 requester: self.instance_id,
863 session_id: self.session_id,
864 },
865 )
866 .await?;
867 }
868
869 Ok(())
870 }
871
872 async fn stage_local_g3_to_g2(&mut self) -> Result<()> {
874 if self.local_g3_blocks.is_empty() {
875 return Ok(());
876 }
877
878 let parallel_worker = self
879 .parallel_worker
880 .as_ref()
881 .ok_or_else(|| anyhow::anyhow!("ParallelWorker required for G3→G2 staging"))?;
882
883 let result =
884 staging::stage_g3_to_g2(&self.local_g3_blocks, &self.g2_manager, &**parallel_worker)
885 .await?;
886
887 let _ = self.local_g3_blocks.take_all();
888 self.local_g2_blocks.extend(result.new_g2_blocks);
889
890 Ok(())
891 }
892
893 async fn pull_remote_blocks(&mut self) -> Result<()> {
901 let parallel_worker = self
902 .parallel_worker
903 .as_ref()
904 .ok_or_else(|| anyhow::anyhow!("ParallelWorker required for RDMA pull"))?;
905
906 for (remote_instance, block_ids) in self.remote_g2_blocks.clone() {
908 if block_ids.is_empty() {
910 continue;
911 }
912
913 let seq_hashes = self
915 .remote_g2_hashes
916 .get(&remote_instance)
917 .cloned()
918 .unwrap_or_default();
919 if seq_hashes.len() != block_ids.len() {
920 anyhow::bail!(
921 "Mismatch between block_ids ({}) and seq_hashes ({}) for instance {}",
922 block_ids.len(),
923 seq_hashes.len(),
924 remote_instance
925 );
926 }
927
928 let mut pairs: Vec<(BlockId, SequenceHash)> =
931 block_ids.into_iter().zip(seq_hashes.into_iter()).collect();
932 pairs.sort_by_key(|(_, hash)| hash.position());
933
934 let block_ids: Vec<BlockId> = pairs.iter().map(|(id, _)| *id).collect();
935 let seq_hashes: Vec<SequenceHash> = pairs.iter().map(|(_, hash)| *hash).collect();
936
937 if !parallel_worker.has_remote_metadata(remote_instance) {
939 tracing::debug!(
940 session_id = %self.session_id,
941 instance = %remote_instance,
942 "Requesting metadata from instance"
943 );
944 let metadata = self.transport.request_metadata(remote_instance).await?;
945 parallel_worker
946 .connect_remote(remote_instance, metadata)?
947 .await?;
948 tracing::debug!(
949 session_id = %self.session_id,
950 instance = %remote_instance,
951 "Metadata imported for instance"
952 );
953 }
954
955 let dst_blocks = self
957 .g2_manager
958 .allocate_blocks(block_ids.len())
959 .ok_or_else(|| {
960 anyhow::anyhow!("Failed to allocate {} G2 blocks", block_ids.len())
961 })?;
962 let dst_ids: Vec<BlockId> = dst_blocks.iter().map(|b| b.block_id()).collect();
963
964 tracing::debug!(
965 session_id = %self.session_id,
966 count = block_ids.len(),
967 instance = %remote_instance,
968 "Pulling blocks via RDMA"
969 );
970
971 let notification = parallel_worker.execute_remote_onboard_for_instance(
974 remote_instance,
975 LogicalLayoutHandle::G2, block_ids,
977 LogicalLayoutHandle::G2, Arc::from(dst_ids),
979 TransferOptions::default(),
980 )?;
981 notification.await?;
982
983 tracing::debug!(
984 session_id = %self.session_id,
985 instance = %remote_instance,
986 "RDMA transfer complete"
987 );
988
989 let new_g2_blocks: Vec<ImmutableBlock<G2>> = dst_blocks
993 .into_iter()
994 .zip(seq_hashes.iter())
995 .map(|(dst, seq_hash)| {
996 let complete = dst
997 .stage(*seq_hash, self.g2_manager.block_size())
998 .expect("block size mismatch");
999 self.g2_manager.register_block(complete)
1000 })
1001 .collect();
1002
1003 self.local_g2_blocks.extend(new_g2_blocks);
1005 }
1006
1007 Ok(())
1008 }
1009
1010 async fn consolidate_blocks(&mut self) {
1017 let mut all_blocks = self.local_g2_blocks.take_all();
1018
1019 all_blocks.sort_by_key(|b| b.sequence_hash().position());
1022
1023 let seq_hashes: Vec<SequenceHash> = all_blocks.iter().map(|b| b.sequence_hash()).collect();
1030 if let Err(e) = validate_contiguous_positions(&seq_hashes) {
1031 tracing::warn!(
1032 session_id = %self.session_id,
1033 error = %e,
1034 "Block positions are not contiguous — proceeding with sorted order"
1035 );
1036 }
1037
1038 let matched_blocks = all_blocks.len();
1039 *self.all_g2_blocks.lock().await = Some(all_blocks);
1040
1041 self.status_tx
1042 .send(OnboardingStatus::Complete { matched_blocks })
1043 .ok();
1044 }
1045
1046 async fn await_commands(&mut self, mut rx: mpsc::Receiver<OnboardMessage>) -> Result<()> {
1048 loop {
1049 tokio::select! {
1050 Some(cmd) = self.control_rx.recv() => {
1051 match cmd {
1052 SessionControl::Prepare => {
1053 if self.mode == StagingMode::Hold {
1054 self.prepare_mode(&mut rx).await?;
1055 self.mode = StagingMode::Prepare;
1056 }
1057 }
1058 SessionControl::Pull => {
1059 if self.mode == StagingMode::Prepare {
1060 self.pull_remote_blocks().await?;
1061 self.consolidate_blocks().await;
1062
1063 let all_remotes: HashSet<InstanceId> = self
1065 .remote_g2_blocks
1066 .keys()
1067 .chain(self.remote_g3_blocks.keys())
1068 .copied()
1069 .collect();
1070
1071 for remote in all_remotes {
1072 self.transport.send(remote, OnboardMessage::CloseSession {
1073 requester: self.instance_id,
1074 session_id: self.session_id,
1075 }).await?;
1076 }
1077
1078 break;
1079 }
1080 }
1081 SessionControl::Cancel => {
1082 let all_remotes: HashSet<InstanceId> = self
1084 .remote_g2_blocks
1085 .keys()
1086 .chain(self.remote_g3_blocks.keys())
1087 .copied()
1088 .collect();
1089
1090 for remote in all_remotes {
1091 self.transport.send(remote, OnboardMessage::CloseSession {
1092 requester: self.instance_id,
1093 session_id: self.session_id,
1094 }).await?;
1095 }
1096 break;
1097 }
1098 SessionControl::Shutdown => {
1099 break;
1100 }
1101 }
1102 }
1103 Some(_msg) = rx.recv() => {
1105 }
1107 }
1108 }
1109
1110 Ok(())
1111 }
1112
1113 fn spawn_g4_search(
1122 &self,
1123 sequence_hashes: Vec<SequenceHash>,
1124 tx: mpsc::Sender<OnboardMessage>,
1125 ) -> JoinHandle<()> {
1126 let session_id = self.session_id;
1127 let parallel_worker = self.parallel_worker.clone();
1129
1130 tokio::spawn(async move {
1131 let Some(worker) = parallel_worker else {
1132 let _ = tx
1134 .send(OnboardMessage::G4Results {
1135 session_id,
1136 found_hashes: vec![],
1137 })
1138 .await;
1139 return;
1140 };
1141
1142 let results = worker.has_blocks(sequence_hashes).await;
1144
1145 let found_hashes: Vec<(SequenceHash, usize)> = results
1147 .into_iter()
1148 .filter_map(|(hash, size_opt)| size_opt.map(|size| (hash, size)))
1149 .collect();
1150
1151 tracing::debug!(
1152 session_id = %session_id,
1153 count = found_hashes.len(),
1154 "G4 search: found blocks in object storage"
1155 );
1156
1157 let _ = tx
1159 .send(OnboardMessage::G4Results {
1160 session_id,
1161 found_hashes,
1162 })
1163 .await;
1164 })
1165 }
1166
1167 fn process_g4_results(
1171 &mut self,
1172 found_hashes: Vec<(SequenceHash, usize)>,
1173 matched_hashes: &mut HashSet<SequenceHash>,
1174 ) -> Vec<SequenceHash> {
1175 let mut won_hashes = Vec::new();
1176
1177 for (hash, _size) in found_hashes {
1178 if matched_hashes.insert(hash) {
1180 won_hashes.push(hash);
1181 self.g4_state.won_hashes.insert(hash);
1182 }
1183 }
1184
1185 tracing::debug!(
1186 session_id = %self.session_id,
1187 won_count = won_hashes.len(),
1188 "G4 won hashes (first-responder-wins)"
1189 );
1190
1191 won_hashes
1192 }
1193
1194 async fn load_g4_blocks(
1200 &mut self,
1201 won_hashes: Vec<SequenceHash>,
1202 g4_tx: mpsc::Sender<OnboardMessage>,
1203 ) -> Result<()> {
1204 if won_hashes.is_empty() {
1205 return Ok(());
1206 }
1207
1208 let parallel_worker = self
1209 .parallel_worker
1210 .as_ref()
1211 .ok_or_else(|| anyhow::anyhow!("ParallelWorkers required for G4 load"))?;
1212
1213 for hash in &won_hashes {
1215 self.g4_state.pending_load.insert(*hash);
1216 }
1217
1218 let dst_blocks = self
1220 .g2_manager
1221 .allocate_blocks(won_hashes.len())
1222 .ok_or_else(|| {
1223 anyhow::anyhow!(
1224 "Failed to allocate {} G2 blocks for G4 load",
1225 won_hashes.len()
1226 )
1227 })?;
1228
1229 let dst_ids: Vec<BlockId> = dst_blocks.iter().map(|b| b.block_id()).collect();
1230
1231 for (hash, block_id) in won_hashes.iter().zip(dst_ids.iter()) {
1233 self.g4_state.allocated_blocks.insert(*hash, *block_id);
1234 }
1235
1236 tracing::debug!(
1237 session_id = %self.session_id,
1238 count = won_hashes.len(),
1239 "Loading G4 blocks via workers"
1240 );
1241
1242 let session_id = self.session_id;
1244 let hashes = won_hashes.clone();
1245 let parallel_worker = parallel_worker.clone();
1246 let g2_manager = self.g2_manager.clone();
1247
1248 tokio::spawn(async move {
1251 let results = parallel_worker
1253 .get_blocks(hashes.clone(), LogicalLayoutHandle::G2, dst_ids.clone())
1254 .await;
1255
1256 let mut success = Vec::new();
1258 let mut failures = Vec::new();
1259 let mut blocks = Vec::new();
1260
1261 for ((result, dst_block), seq_hash) in results
1263 .into_iter()
1264 .zip(dst_blocks.into_iter())
1265 .zip(hashes.iter())
1266 {
1267 match result {
1268 Ok(hash) => {
1269 let complete = dst_block
1272 .stage(*seq_hash, g2_manager.block_size())
1273 .expect("block size mismatch");
1274 let immutable = g2_manager.register_block(complete);
1275 blocks.push(immutable);
1276 success.push(hash);
1277 }
1278 Err(hash) => {
1279 failures.push((hash, "Failed to download block".to_string()));
1281 }
1282 }
1283 }
1284
1285 tracing::debug!(
1286 session_id = %session_id,
1287 success_count = success.len(),
1288 failure_count = failures.len(),
1289 "G4 load complete"
1290 );
1291
1292 let _ = g4_tx
1294 .send(OnboardMessage::G4LoadComplete {
1295 session_id,
1296 success,
1297 failures,
1298 blocks: std::sync::Arc::new(blocks),
1299 })
1300 .await;
1301 });
1302
1303 Ok(())
1304 }
1305
1306 fn handle_g4_load_complete(
1311 &mut self,
1312 success: Vec<SequenceHash>,
1313 failures: Vec<(SequenceHash, String)>,
1314 blocks: Arc<Vec<ImmutableBlock<G2>>>,
1315 ) {
1316 for hash in &success {
1318 self.g4_state.pending_load.remove(hash);
1319 self.g4_state.allocated_blocks.remove(hash);
1321 }
1322
1323 let blocks =
1325 Arc::try_unwrap(blocks).expect("G4LoadComplete should be the sole owner of blocks");
1326
1327 self.local_g2_blocks.extend(blocks);
1331
1332 for (hash, error) in failures {
1334 self.g4_state.pending_load.remove(&hash);
1335 self.g4_state.failed_hashes.insert(hash, error);
1336
1337 self.g4_state.allocated_blocks.remove(&hash);
1339
1340 self.g4_state.won_hashes.remove(&hash);
1342 }
1343
1344 tracing::debug!(
1345 session_id = %self.session_id,
1346 won = self.g4_state.won_hashes.len(),
1347 pending = self.g4_state.pending_load.len(),
1348 failed = self.g4_state.failed_hashes.len(),
1349 local_g2 = self.local_g2_blocks.count(),
1350 "G4 load complete, blocks added to local_g2_blocks"
1351 );
1352 }
1353}