use std::sync::Arc;
use parking_lot::RwLock;
use rustc_hash::{FxHashMap, FxHashSet};
use std::collections::VecDeque;
use super::{
EventKind, EventWarningKind, KvIndexerMetrics, PreBoundEventCounters, SyncIndexer,
WorkerLookupStats, WorkerTask,
};
use crate::active_set::reconcile_active_workers;
use crate::cleanup::{self, CleanableNode, CleanupGuard, CleanupState};
use crate::protocols::*;
type SharedBlock = Arc<RwLock<Block>>;
type WorkerLookup = FxHashMap<ExternalSequenceBlockHash, SharedBlock>;
#[derive(Debug)]
struct Block {
children: FxHashMap<LocalBlockHash, SharedBlock>,
workers: FxHashSet<WorkerWithDpRank>,
block_hash: Option<ExternalSequenceBlockHash>,
}
impl Block {
fn new() -> Self {
Self {
children: FxHashMap::default(),
workers: FxHashSet::default(),
block_hash: None,
}
}
fn with_hash(block_hash: ExternalSequenceBlockHash) -> Self {
Self {
children: FxHashMap::default(),
workers: FxHashSet::default(),
block_hash: Some(block_hash),
}
}
#[inline]
fn drop_worker(&mut self, worker: WorkerWithDpRank) {
self.workers.remove(&worker);
if self.workers.is_empty() {
self.children.clear();
}
}
}
impl CleanableNode for Block {
type ChildKey = LocalBlockHash;
fn has_any_workers(&self) -> bool {
!self.workers.is_empty()
}
fn children(&self) -> &FxHashMap<LocalBlockHash, SharedBlock> {
&self.children
}
fn remove_child(&mut self, key: &LocalBlockHash) {
self.children.remove(key);
}
}
pub struct ConcurrentRadixTree {
root: SharedBlock,
cleanup: CleanupState,
}
impl Default for ConcurrentRadixTree {
fn default() -> Self {
Self::new()
}
}
impl Drop for ConcurrentRadixTree {
fn drop(&mut self) {
let mut stack: Vec<SharedBlock> = Vec::new();
{
let mut root = self.root.write();
stack.extend(root.children.drain().map(|(_, v)| v));
}
while let Some(block) = stack.pop() {
if let Ok(rwlock) = Arc::try_unwrap(block) {
let mut inner = rwlock.into_inner();
stack.extend(inner.children.drain().map(|(_, v)| v));
}
}
}
}
impl ConcurrentRadixTree {
pub fn new() -> Self {
Self {
root: Arc::new(RwLock::new(Block::new())),
cleanup: CleanupState::new(),
}
}
pub fn find_matches_impl(
&self,
sequence: &[LocalBlockHash],
early_exit: bool,
) -> OverlapScores {
let mut scores = OverlapScores::new();
if sequence.is_empty() {
return scores;
}
let first_child = {
let guard = self.root.read();
guard.children.get(&sequence[0]).cloned()
};
let Some(first_child) = first_child else {
return scores;
};
let (mut active, mut active_count) = {
let guard = first_child.read();
(guard.workers.clone(), guard.workers.len())
};
if active.is_empty() {
return scores;
}
if early_exit && active_count == 1 {
for worker in &active {
scores.scores.insert(*worker, 1);
}
return scores;
}
let mut current = first_child;
let mut matched_depth = 1u32;
for (idx, local_hash) in sequence.iter().enumerate().skip(1) {
let next_block = {
let guard = current.read();
guard.children.get(local_hash).cloned()
};
let Some(block) = next_block else {
break;
};
{
let guard = block.read();
let child_count = guard.workers.len();
if child_count != active_count {
reconcile_active_workers(&mut active, &guard.workers, |worker| {
scores.scores.insert(worker, matched_depth);
});
active_count = active.len();
if active_count == 0 {
break;
}
}
if early_exit && active_count == 1 {
matched_depth = (idx + 1) as u32;
break;
}
}
current = block;
matched_depth = (idx + 1) as u32;
}
for worker in &active {
scores.scores.insert(*worker, matched_depth);
}
scores
}
fn apply_event(
&self,
lookup: &mut FxHashMap<WorkerWithDpRank, WorkerLookup>,
event: RouterEvent,
counters: Option<&PreBoundEventCounters>,
) -> Result<(), KvCacheEventError> {
let (worker_id, kv_event) = (event.worker_id, event.event);
let (id, op) = (kv_event.event_id, kv_event.data);
let worker = WorkerWithDpRank::new(worker_id, kv_event.dp_rank);
match op {
KvCacheEventData::Stored(op) => self.apply_stored(lookup, worker, op, id, counters),
KvCacheEventData::Removed(op) => self.apply_removed(lookup, worker, op, id),
KvCacheEventData::Cleared => {
lookup.entry(worker).or_default();
self.clear_all_blocks(lookup, worker.worker_id);
Ok(())
}
}
}
fn apply_stored(
&self,
lookup: &mut FxHashMap<WorkerWithDpRank, WorkerLookup>,
worker: WorkerWithDpRank,
op: KvCacheStoreData,
id: u64,
counters: Option<&PreBoundEventCounters>,
) -> Result<(), KvCacheEventError> {
let worker_lookup = lookup.entry(worker).or_default();
let mut current = match op.parent_hash {
Some(parent) => match worker_lookup.get(&parent) {
Some(block) => block.clone(),
None => {
tracing::warn!(
worker_id = worker.worker_id.to_string(),
dp_rank = worker.dp_rank,
id,
parent_hash = ?op.parent_hash,
num_blocks = op.blocks.len(),
"Failed to find parent block; skipping store operation"
);
return Err(KvCacheEventError::ParentBlockNotFound);
}
},
None => self.root.clone(),
};
let mut needs_worker_insert = false;
let mut duplicate_store = !op.blocks.is_empty();
for block_data in op.blocks {
let child = {
let mut parent_guard = current.write();
if needs_worker_insert && parent_guard.workers.insert(worker) {
duplicate_store = false;
}
needs_worker_insert = true;
match parent_guard.children.get(&block_data.tokens_hash) {
Some(existing) => {
{
let existing_guard = existing.read();
if existing_guard.block_hash != Some(block_data.block_hash) {
duplicate_store = false;
tracing::warn!(
expected = ?block_data.block_hash,
actual = ?existing_guard.block_hash,
"block_hash mismatch: sequence hashes should be uniform across workers"
);
}
}
existing.clone()
}
None => {
duplicate_store = false;
let new_block = worker_lookup
.get(&block_data.block_hash)
.cloned()
.unwrap_or_else(|| {
Arc::new(RwLock::new(Block::with_hash(block_data.block_hash)))
});
parent_guard
.children
.insert(block_data.tokens_hash, new_block.clone());
new_block
}
}
};
match worker_lookup.insert(block_data.block_hash, child.clone()) {
Some(existing) if Arc::ptr_eq(&existing, &child) => {}
Some(_) => duplicate_store = false,
None => {
duplicate_store = false;
}
}
current = child;
}
if needs_worker_insert && current.write().workers.insert(worker) {
duplicate_store = false;
}
if duplicate_store && let Some(counters) = counters {
counters.inc_warning(EventWarningKind::DuplicateStore);
}
Ok(())
}
fn apply_removed(
&self,
lookup: &mut FxHashMap<WorkerWithDpRank, WorkerLookup>,
worker: WorkerWithDpRank,
op: KvCacheRemoveData,
id: u64,
) -> Result<(), KvCacheEventError> {
let Some(worker_lookup) = lookup.get_mut(&worker) else {
return Err(KvCacheEventError::BlockNotFound);
};
for block_hash in op.block_hashes {
let Some(block) = worker_lookup.remove(&block_hash) else {
tracing::debug!(
worker_id = worker.worker_id.to_string(),
dp_rank = worker.dp_rank,
id,
block_hash = ?block_hash,
"Block not found during remove; skipping"
);
continue;
};
block.write().drop_worker(worker);
}
Ok(())
}
fn remove_or_clear_worker_blocks(
&self,
lookup: &mut FxHashMap<WorkerWithDpRank, WorkerLookup>,
worker_id: WorkerId,
keep_worker: bool,
) {
let workers: Vec<WorkerWithDpRank> = lookup
.keys()
.filter(|w| w.worker_id == worker_id)
.copied()
.collect();
for worker in workers {
if let Some(worker_lookup) = lookup.remove(&worker) {
for (_, block) in worker_lookup.into_iter() {
block.write().drop_worker(worker);
}
if keep_worker {
lookup.insert(worker, FxHashMap::default());
}
}
}
}
fn remove_worker_dp_rank(
&self,
lookup: &mut FxHashMap<WorkerWithDpRank, WorkerLookup>,
worker_id: WorkerId,
dp_rank: DpRank,
) {
let key = WorkerWithDpRank { worker_id, dp_rank };
if let Some(worker_lookup) = lookup.remove(&key) {
for (_, block) in worker_lookup.into_iter() {
block.write().drop_worker(key);
}
}
}
fn clear_all_blocks(
&self,
lookup: &mut FxHashMap<WorkerWithDpRank, WorkerLookup>,
worker_id: WorkerId,
) {
self.remove_or_clear_worker_blocks(lookup, worker_id, true);
}
fn dump_tree_as_events(&self) -> Vec<RouterEvent> {
tracing::debug!("Dumping concurrent radix tree as events");
let mut events = Vec::new();
let mut event_id = 0u64;
let mut queue = VecDeque::new();
{
let root_guard = self.root.read();
for (tokens_hash, child_block) in &root_guard.children {
queue.push_back((child_block.clone(), None, *tokens_hash));
}
}
while let Some((current_block, parent_hash, tokens_hash)) = queue.pop_front() {
let current_guard = current_block.read();
let block_hash = current_guard
.block_hash
.expect("non-root block must have block_hash");
for worker in ¤t_guard.workers {
let event = RouterEvent {
worker_id: worker.worker_id,
storage_tier: crate::protocols::StorageTier::Device,
event: KvCacheEvent {
event_id,
data: KvCacheEventData::Stored(KvCacheStoreData {
parent_hash,
start_position: None,
blocks: vec![KvCacheStoredBlockData {
block_hash,
mm_extra_info: None,
tokens_hash,
}],
}),
dp_rank: worker.dp_rank,
},
};
events.push(event);
event_id += 1;
}
for (child_tokens_hash, child_block) in ¤t_guard.children {
queue.push_back((child_block.clone(), Some(block_hash), *child_tokens_hash));
}
}
events
}
}
impl SyncIndexer for ConcurrentRadixTree {
fn worker(
&self,
event_receiver: flume::Receiver<WorkerTask>,
metrics: Option<Arc<KvIndexerMetrics>>,
) -> anyhow::Result<()> {
let mut lookup = FxHashMap::default();
let counters = metrics.as_ref().map(|m| m.prebind());
while let Ok(task) = event_receiver.recv() {
match task {
WorkerTask::Event(event) => {
let kind = EventKind::of(&event.event.data);
let result = self.apply_event(&mut lookup, event, counters.as_ref());
if result.is_err() {
tracing::warn!("Failed to apply event: {:?}", result.as_ref().err());
}
if let Some(ref c) = counters {
c.inc(kind, result);
}
}
WorkerTask::EventWithAck { event, resp } => {
let kind = EventKind::of(&event.event.data);
let result = self.apply_event(&mut lookup, event, counters.as_ref());
let applied = result.is_ok();
if result.is_err() {
tracing::warn!("Failed to apply event: {:?}", result.as_ref().err());
}
if let Some(ref c) = counters {
c.inc(kind, result);
}
let _ = resp.send(applied);
}
WorkerTask::Anchor { worker, anchor } => {
if let Err(error) = self.apply_anchor(worker, anchor) {
tracing::warn!(?error, "Failed to apply anchor");
}
}
WorkerTask::RemoveWorker(worker_id) => {
self.remove_or_clear_worker_blocks(&mut lookup, worker_id, false);
}
WorkerTask::RemoveWorkerDpRank(worker_id, dp_rank) => {
self.remove_worker_dp_rank(&mut lookup, worker_id, dp_rank);
}
WorkerTask::CleanupStaleChildren => {
self.run_cleanup_task();
}
WorkerTask::DumpEvents(_sender) => {
let _ = _sender.send(Ok(Vec::new()));
}
WorkerTask::Stats(sender) => {
let stats = WorkerLookupStats::from_worker_block_counts(
lookup
.iter()
.map(|(worker, worker_lookup)| (*worker, worker_lookup.len())),
);
let _ = sender.send(stats);
}
WorkerTask::Flush(sender) => {
let _ = sender.send(());
}
WorkerTask::Terminate => {
break;
}
}
}
tracing::debug!("ConcurrentRadixTree worker thread shutting down");
Ok(())
}
fn find_matches(&self, sequence: &[LocalBlockHash], early_exit: bool) -> OverlapScores {
self.find_matches_impl(sequence, early_exit)
}
fn try_schedule_cleanup(&self) -> bool {
self.cleanup.try_schedule()
}
fn cancel_scheduled_cleanup(&self) {
self.cleanup.cancel();
}
fn run_cleanup_task(&self) {
let mut cleanup_guard = CleanupGuard::new(&self.cleanup);
cleanup::sweep_stale_children(&self.root);
cleanup_guard.mark_completed();
}
fn dump_events(&self) -> Option<Vec<RouterEvent>> {
Some(self.dump_tree_as_events())
}
}