use anyhow::Result;
use tokio::sync::mpsc;
use std::collections::HashSet;
use std::sync::Arc;
use crate::{BlockId, G2, G3, InstanceId, SequenceHash, worker::group::ParallelWorkers};
use kvbm_logical::manager::BlockManager;
use super::{BlockHolder, SessionId, messages::OnboardMessage, transport::MessageTransport};
pub struct ResponderSession {
session_id: SessionId,
instance_id: InstanceId,
requester: InstanceId,
g2_manager: Arc<BlockManager<G2>>,
g3_manager: Option<Arc<BlockManager<G3>>>,
parallel_worker: Option<Arc<dyn ParallelWorkers>>,
transport: Arc<MessageTransport>,
held_g2_blocks: BlockHolder<G2>,
held_g3_blocks: BlockHolder<G3>,
}
impl ResponderSession {
pub fn new(
session_id: SessionId,
instance_id: InstanceId,
requester: InstanceId,
g2_manager: Arc<BlockManager<G2>>,
g3_manager: Option<Arc<BlockManager<G3>>>,
parallel_worker: Option<Arc<dyn ParallelWorkers>>,
transport: Arc<MessageTransport>,
) -> Self {
Self {
session_id,
instance_id,
requester,
g2_manager,
g3_manager,
parallel_worker,
transport,
held_g2_blocks: BlockHolder::empty(),
held_g3_blocks: BlockHolder::empty(),
}
}
pub async fn run(
mut self,
mut rx: mpsc::Receiver<OnboardMessage>,
sequence_hashes: Vec<SequenceHash>,
) -> Result<()> {
let g2_matches_map = self.g2_manager.scan_matches(&sequence_hashes, true);
let mut g2_matches: Vec<_> = g2_matches_map.into_values().collect();
g2_matches.sort_by_key(|block| block.sequence_hash().position());
self.held_g2_blocks = BlockHolder::new(g2_matches);
let g2_sequence_hashes: Vec<SequenceHash> = self.held_g2_blocks.sequence_hashes();
let g2_block_ids: Vec<BlockId> = self
.held_g2_blocks
.blocks()
.iter()
.map(|b| b.block_id())
.collect();
let g2_msg = OnboardMessage::G2Results {
responder: self.instance_id,
session_id: self.session_id,
sequence_hashes: g2_sequence_hashes,
block_ids: g2_block_ids,
};
self.transport.send(self.requester, g2_msg).await?;
let g2_matched_hashes: HashSet<SequenceHash> =
self.held_g2_blocks.sequence_hashes().into_iter().collect();
let remaining_hashes: Vec<SequenceHash> = sequence_hashes
.iter()
.filter(|h| !g2_matched_hashes.contains(h))
.copied()
.collect();
if !remaining_hashes.is_empty()
&& let Some(ref g3_manager) = self.g3_manager
{
let g3_matches_map = g3_manager.scan_matches(&remaining_hashes, true);
let mut g3_matches: Vec<_> = g3_matches_map.into_values().collect();
g3_matches.sort_by_key(|block| block.sequence_hash().position());
if !g3_matches.is_empty() {
self.held_g3_blocks = BlockHolder::new(g3_matches);
let g3_sequence_hashes: Vec<SequenceHash> = self.held_g3_blocks.sequence_hashes();
let g3_msg = OnboardMessage::G3Results {
responder: self.instance_id,
session_id: self.session_id,
sequence_hashes: g3_sequence_hashes,
};
self.transport.send(self.requester, g3_msg).await?;
}
}
let complete_msg = OnboardMessage::SearchComplete {
responder: self.instance_id,
session_id: self.session_id,
};
self.transport.send(self.requester, complete_msg).await?;
while let Some(msg) = rx.recv().await {
match msg {
OnboardMessage::HoldBlocks {
hold_hashes,
drop_hashes: _,
..
} => {
self.held_g2_blocks.retain(&hold_hashes);
self.held_g3_blocks.retain(&hold_hashes);
let ack = OnboardMessage::Acknowledged {
responder: self.instance_id,
session_id: self.session_id,
};
self.transport.send(self.requester, ack).await?;
}
OnboardMessage::StageBlocks { stage_hashes, .. } => {
self.held_g3_blocks.retain(&stage_hashes);
if !self.held_g3_blocks.is_empty() {
if self.parallel_worker.is_some() {
self.stage_g3_to_g2().await?;
} else {
tracing::warn!(
session_id = %self.session_id,
g3_blocks = self.held_g3_blocks.count(),
"G3 blocks cannot be staged: no parallel worker configured"
);
}
}
}
OnboardMessage::ReleaseBlocks { release_hashes, .. } => {
self.held_g2_blocks.release(&release_hashes);
self.held_g3_blocks.release(&release_hashes);
}
OnboardMessage::CloseSession { .. } => {
let _ = self.held_g2_blocks.take_all();
let _ = self.held_g3_blocks.take_all();
break;
}
OnboardMessage::CreateSession { .. } => {
}
_ => {
tracing::warn!(
session_id = %self.session_id,
msg = ?msg,
"ResponderSession: unexpected message"
);
}
}
}
Ok(())
}
async fn stage_g3_to_g2(&mut self) -> Result<()> {
let parallel_worker = self
.parallel_worker
.as_ref()
.ok_or_else(|| anyhow::anyhow!("ParallelWorker required for G3->G2 staging"))?;
let result = super::staging::stage_g3_to_g2(
&self.held_g3_blocks,
&self.g2_manager,
&**parallel_worker,
)
.await?;
let new_sequence_hashes: Vec<SequenceHash> = result
.new_g2_blocks
.iter()
.map(|b| b.sequence_hash())
.collect();
let new_block_ids: Vec<BlockId> =
result.new_g2_blocks.iter().map(|b| b.block_id()).collect();
let _ = self.held_g3_blocks.take_all();
self.held_g2_blocks.extend(result.new_g2_blocks);
let ready_msg = OnboardMessage::BlocksReady {
responder: self.instance_id,
session_id: self.session_id,
sequence_hashes: new_sequence_hashes,
block_ids: new_block_ids,
};
self.transport.send(self.requester, ready_msg).await?;
Ok(())
}
}