use dashmap::DashMap;
use dashmap::mapref::entry::Entry;
use rustc_hash::{FxBuildHasher, FxHashMap, FxHashSet};
use std::sync::Arc;
use super::{
EventKind, EventWarningKind, KvIndexerMetrics, PreBoundEventCounters, SyncIndexer,
WorkerLookupStats, WorkerTask,
};
use crate::active_set::reconcile_active_workers;
use crate::protocols::{
DpRank, ExternalSequenceBlockHash, KvCacheEvent, KvCacheEventData, KvCacheEventError,
KvCacheStoreData, KvCacheStoredBlockData, LocalBlockHash, OverlapScores, RouterEvent, WorkerId,
WorkerWithDpRank,
};
pub const DYN_ROUTER_POSITIONAL_SEARCH_MODE: &str = "DYN_ROUTER_POSITIONAL_SEARCH_MODE";
#[derive(Debug, Clone, Copy, PartialEq, Eq, Default)]
pub enum SearchMode {
#[default]
Strided,
Binary,
}
impl SearchMode {
fn from_env() -> Self {
match std::env::var(DYN_ROUTER_POSITIONAL_SEARCH_MODE) {
Ok(value) => match value.trim().to_ascii_lowercase().as_str() {
"strided" => SearchMode::Strided,
"binary" => SearchMode::Binary,
other => {
tracing::warn!(
value = %other,
"invalid {DYN_ROUTER_POSITIONAL_SEARCH_MODE}, expected 'strided' or 'binary'; falling back to strided"
);
SearchMode::Strided
}
},
Err(_) => SearchMode::Strided,
}
}
}
#[derive(Debug, Clone)]
enum SeqEntry {
Single(ExternalSequenceBlockHash, FxHashSet<WorkerWithDpRank>),
Multi(FxHashMap<ExternalSequenceBlockHash, FxHashSet<WorkerWithDpRank>>),
}
impl SeqEntry {
fn new(seq_hash: ExternalSequenceBlockHash, worker: WorkerWithDpRank) -> Self {
let mut workers = FxHashSet::default();
workers.insert(worker);
Self::Single(seq_hash, workers)
}
fn insert(&mut self, seq_hash: ExternalSequenceBlockHash, worker: WorkerWithDpRank) -> bool {
match self {
Self::Single(existing_hash, workers) if *existing_hash == seq_hash => {
workers.insert(worker)
}
Self::Single(existing_hash, existing_workers) => {
let mut map = FxHashMap::with_capacity_and_hasher(2, FxBuildHasher);
map.insert(*existing_hash, std::mem::take(existing_workers));
map.entry(seq_hash).or_default().insert(worker);
*self = Self::Multi(map);
true
}
Self::Multi(map) => map.entry(seq_hash).or_default().insert(worker),
}
}
fn remove(&mut self, seq_hash: ExternalSequenceBlockHash, worker: WorkerWithDpRank) -> bool {
match self {
Self::Single(existing_hash, workers) if *existing_hash == seq_hash => {
workers.remove(&worker);
workers.is_empty()
}
Self::Single(_, _) => false, Self::Multi(map) => {
if let Some(workers) = map.get_mut(&seq_hash) {
workers.remove(&worker);
if workers.is_empty() {
map.remove(&seq_hash);
}
}
map.is_empty()
}
}
}
fn get(&self, seq_hash: ExternalSequenceBlockHash) -> Option<&FxHashSet<WorkerWithDpRank>> {
match self {
Self::Single(existing_hash, workers) if *existing_hash == seq_hash => Some(workers),
Self::Single(_, _) => None,
Self::Multi(map) => map.get(&seq_hash),
}
}
}
pub type LevelIndex = FxHashMap<ExternalSequenceBlockHash, (usize, LocalBlockHash)>;
pub struct PositionalIndexer {
index: DashMap<(usize, LocalBlockHash), SeqEntry, FxBuildHasher>,
jump_size: usize,
search_mode: SearchMode,
}
impl PositionalIndexer {
pub fn new(jump_size: usize) -> Self {
Self::new_with_mode(jump_size, SearchMode::from_env())
}
pub fn new_with_mode(jump_size: usize, search_mode: SearchMode) -> Self {
assert!(jump_size > 0, "jump_size must be greater than 0");
Self {
index: DashMap::with_hasher(FxBuildHasher),
jump_size,
search_mode,
}
}
pub fn with_search_mode(mut self, search_mode: SearchMode) -> Self {
self.search_mode = search_mode;
self
}
}
impl SyncIndexer for PositionalIndexer {
fn worker(
&self,
event_receiver: flume::Receiver<WorkerTask>,
metrics: Option<Arc<KvIndexerMetrics>>,
) -> anyhow::Result<()> {
let mut worker_blocks = 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 worker_blocks, 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 worker_blocks, 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_impl(&mut worker_blocks, worker_id, false);
}
WorkerTask::RemoveWorkerDpRank(worker_id, dp_rank) => {
self.remove_worker_dp_rank_impl(&mut worker_blocks, worker_id, dp_rank);
}
WorkerTask::CleanupStaleChildren => {
self.run_cleanup_task();
}
WorkerTask::DumpEvents(sender) => {
let events = self.dump_events(&worker_blocks);
if let Err(e) = sender.send(Ok(events)) {
tracing::warn!("Failed to send events: {:?}", e);
}
}
WorkerTask::Stats(sender) => {
let stats = WorkerLookupStats::from_worker_block_counts(
worker_blocks
.iter()
.map(|(worker, worker_map)| (*worker, worker_map.len())),
);
let _ = sender.send(stats);
}
WorkerTask::Flush(sender) => {
let _ = sender.send(());
}
WorkerTask::Terminate => {
break;
}
}
}
tracing::debug!("PositionalIndexer worker thread shutting down");
Ok(())
}
fn find_matches(&self, sequence: &[LocalBlockHash], early_exit: bool) -> OverlapScores {
match self.search_mode {
SearchMode::Strided => self.jump_search_matches(sequence, early_exit),
SearchMode::Binary => self.binary_search_matches(sequence, early_exit),
}
}
}
impl PositionalIndexer {
pub fn apply_event(
&self,
worker_blocks: &mut FxHashMap<WorkerWithDpRank, LevelIndex>,
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);
tracing::trace!(
id,
"PositionalIndexer::apply_event_impl: operation: {:?}",
op
);
match op {
KvCacheEventData::Stored(store_data) => {
self.store_blocks_impl(worker_blocks, worker, store_data, id, counters)?;
Ok(())
}
KvCacheEventData::Removed(remove_data) => {
self.remove_blocks_impl(worker_blocks, worker, &remove_data.block_hashes, id)?;
Ok(())
}
KvCacheEventData::Cleared => {
self.clear_worker_blocks_impl(worker_blocks, worker_id);
Ok(())
}
}
}
fn store_blocks_impl(
&self,
worker_blocks: &mut FxHashMap<WorkerWithDpRank, LevelIndex>,
worker: WorkerWithDpRank,
store_data: KvCacheStoreData,
event_id: u64,
counters: Option<&PreBoundEventCounters>,
) -> Result<(), KvCacheEventError> {
let KvCacheStoreData {
parent_hash,
start_position,
blocks,
} = store_data;
let worker_map = worker_blocks.entry(worker).or_default();
let start_pos = match start_position {
Some(start_position) => start_position as usize,
None => match parent_hash {
Some(parent_hash) => {
let Some(entry) = worker_map.get(&parent_hash) else {
tracing::warn!(
worker_id = worker.worker_id.to_string(),
dp_rank = worker.dp_rank,
event_id,
parent_hash = ?parent_hash,
);
return Err(KvCacheEventError::ParentBlockNotFound);
};
entry.0 + 1 }
None => 0, },
};
let worker_blocks_entry = worker_blocks.entry(worker).or_default();
let mut duplicate_store = !blocks.is_empty();
for (i, block_data) in blocks.into_iter().enumerate() {
let position = start_pos + i;
let local_hash = block_data.tokens_hash;
let seq_hash = block_data.block_hash;
match self.index.entry((position, local_hash)) {
Entry::Occupied(mut entry) => {
if entry.get_mut().insert(seq_hash, worker) {
duplicate_store = false;
}
}
Entry::Vacant(entry) => {
entry.insert(SeqEntry::new(seq_hash, worker));
duplicate_store = false;
}
}
match worker_blocks_entry.insert(seq_hash, (position, local_hash)) {
Some(existing) if existing == (position, local_hash) => {}
Some(_) => duplicate_store = false,
None => {
duplicate_store = false;
}
}
}
if duplicate_store && let Some(counters) = counters {
counters.inc_warning(EventWarningKind::DuplicateStore);
}
Ok(())
}
fn remove_blocks_impl(
&self,
worker_blocks: &mut FxHashMap<WorkerWithDpRank, LevelIndex>,
worker: WorkerWithDpRank,
seq_hashes: &Vec<ExternalSequenceBlockHash>,
event_id: u64,
) -> Result<(), KvCacheEventError> {
let worker_map = worker_blocks.get_mut(&worker).ok_or_else(|| {
tracing::warn!(
worker_id = worker.worker_id.to_string(),
dp_rank = worker.dp_rank,
event_id,
block_hashes = ?seq_hashes,
"Failed to find worker blocks to remove"
);
KvCacheEventError::BlockNotFound
})?;
for seq_hash in seq_hashes {
let Some((position, local_hash)) = worker_map.remove(seq_hash) else {
tracing::warn!(
worker_id = worker.worker_id.to_string(),
dp_rank = worker.dp_rank,
event_id,
block_hash = ?seq_hash,
"Failed to find block to remove; skipping remove operation"
);
return Err(KvCacheEventError::BlockNotFound);
};
if let Some(mut entry) = self.index.get_mut(&(position, local_hash)) {
let _ = entry.remove(*seq_hash, worker);
}
}
Ok(())
}
fn clear_worker_blocks_impl(
&self,
worker_blocks: &mut FxHashMap<WorkerWithDpRank, LevelIndex>,
worker_id: WorkerId,
) {
self.remove_or_clear_worker_blocks_impl(worker_blocks, worker_id, true);
}
fn remove_worker_dp_rank_impl(
&self,
worker_blocks: &mut FxHashMap<WorkerWithDpRank, LevelIndex>,
worker_id: WorkerId,
dp_rank: DpRank,
) {
let key = WorkerWithDpRank { worker_id, dp_rank };
if let Some(worker_map) = worker_blocks.remove(&key) {
for (seq_hash, (position, local_hash)) in worker_map.iter() {
if let Some(mut entry) = self.index.get_mut(&(*position, *local_hash)) {
let _ = entry.remove(*seq_hash, key);
}
}
}
}
fn remove_or_clear_worker_blocks_impl(
&self,
worker_blocks: &mut FxHashMap<WorkerWithDpRank, LevelIndex>,
worker_id: WorkerId,
keep_worker: bool,
) {
let workers: Vec<WorkerWithDpRank> = worker_blocks
.iter()
.filter(|entry| entry.0.worker_id == worker_id)
.map(|entry| *entry.0)
.collect();
for worker in workers {
if let Some(worker_map) = worker_blocks.remove(&worker) {
for (seq_hash, (position, local_hash)) in worker_map.iter() {
if let Some(mut entry) = self.index.get_mut(&(*position, *local_hash)) {
let _ = entry.remove(*seq_hash, worker);
}
}
}
if keep_worker {
worker_blocks.insert(worker, FxHashMap::default());
}
}
}
fn dump_events(
&self,
worker_blocks: &FxHashMap<WorkerWithDpRank, LevelIndex>,
) -> Vec<RouterEvent> {
let mut events = Vec::new();
let mut event_id = 0u64;
for (worker, worker_map) in worker_blocks.iter() {
let mut blocks: Vec<_> = worker_map
.iter()
.map(|(seq_hash, (pos, local_hash))| (*pos, *local_hash, *seq_hash))
.collect();
blocks.sort_unstable_by_key(|(pos, _, _)| *pos);
for (pos, local_hash, seq_hash) in blocks {
events.push(RouterEvent {
worker_id: worker.worker_id,
storage_tier: crate::protocols::StorageTier::Device,
event: KvCacheEvent {
event_id,
data: KvCacheEventData::Stored(KvCacheStoreData {
parent_hash: None,
start_position: Some(pos as u32),
blocks: vec![KvCacheStoredBlockData {
block_hash: seq_hash,
tokens_hash: local_hash,
mm_extra_info: None,
}],
}),
dp_rank: worker.dp_rank,
},
});
event_id += 1;
}
}
events
}
}
impl PositionalIndexer {
#[inline]
fn compute_next_seq_hash(prev_seq_hash: u64, current_local_hash: u64) -> u64 {
dynamo_tokens::compute_next_sequence_hash(prev_seq_hash, current_local_hash)
}
#[inline]
fn ensure_seq_hash_computed(
seq_hashes: &mut Vec<ExternalSequenceBlockHash>,
target_pos: usize,
sequence: &[LocalBlockHash],
) {
while seq_hashes.len() <= target_pos {
let pos = seq_hashes.len();
if pos == 0 {
seq_hashes.push(ExternalSequenceBlockHash::from(sequence[0].0));
} else {
let prev_seq_hash = seq_hashes[pos - 1].0;
let current_local_hash = sequence[pos].0;
let next_hash = Self::compute_next_seq_hash(prev_seq_hash, current_local_hash);
seq_hashes.push(ExternalSequenceBlockHash::from(next_hash));
}
}
}
fn get_workers_lazy(
&self,
position: usize,
local_hash: LocalBlockHash,
seq_hashes: &mut Vec<ExternalSequenceBlockHash>,
sequence: &[LocalBlockHash],
) -> Option<FxHashSet<WorkerWithDpRank>> {
let entry = self.index.get(&(position, local_hash))?;
Self::ensure_seq_hash_computed(seq_hashes, position, sequence);
let seq_hash = seq_hashes[position];
entry.get(seq_hash).cloned()
}
fn count_workers_at(
&self,
position: usize,
local_hash: LocalBlockHash,
seq_hashes: &mut Vec<ExternalSequenceBlockHash>,
sequence: &[LocalBlockHash],
) -> Option<usize> {
let entry = self.index.get(&(position, local_hash))?;
Self::ensure_seq_hash_computed(seq_hashes, position, sequence);
let seq_hash = seq_hashes[position];
Some(
entry
.get(seq_hash)
.map(|workers| workers.len())
.unwrap_or(0),
)
}
#[expect(clippy::too_many_arguments)]
fn linear_scan_drain(
&self,
sequence: &[LocalBlockHash],
seq_hashes: &mut Vec<ExternalSequenceBlockHash>,
active: &mut FxHashSet<WorkerWithDpRank>,
scores: &mut OverlapScores,
lo: usize,
hi: usize,
early_exit: bool,
) {
if active.is_empty() {
return;
}
for pos in lo..hi {
if active.is_empty() {
break;
}
let Some(entry) = self.index.get(&(pos, sequence[pos])) else {
for worker in active.drain() {
scores.scores.insert(worker, pos as u32);
}
break;
};
Self::ensure_seq_hash_computed(seq_hashes, pos, sequence);
let Some(workers) = entry.get(seq_hashes[pos]) else {
for worker in active.drain() {
scores.scores.insert(worker, pos as u32);
}
break;
};
if workers.len() != active.len() {
reconcile_active_workers(active, workers, |worker| {
scores.scores.insert(worker, pos as u32);
});
}
if early_exit && !active.is_empty() {
break;
}
}
}
fn jump_search_matches(
&self,
local_hashes: &[LocalBlockHash],
early_exit: bool,
) -> OverlapScores {
let mut scores = OverlapScores::new();
if local_hashes.is_empty() {
return scores;
}
let mut seq_hashes: Vec<ExternalSequenceBlockHash> = Vec::with_capacity(local_hashes.len());
let Some(initial_workers) =
self.get_workers_lazy(0, local_hashes[0], &mut seq_hashes, local_hashes)
else {
return scores;
};
let mut active = initial_workers;
if active.is_empty() {
return scores;
}
if early_exit {
for worker in &active {
scores.scores.insert(*worker, 1);
}
return scores;
}
let len = local_hashes.len();
let mut current_pos = 0;
while current_pos < len - 1 && !active.is_empty() {
let next_pos = (current_pos + self.jump_size).min(len - 1);
let num_workers_at_next = self
.count_workers_at(
next_pos,
local_hashes[next_pos],
&mut seq_hashes,
local_hashes,
)
.unwrap_or(0);
if num_workers_at_next == active.len() {
current_pos = next_pos;
} else {
self.linear_scan_drain(
local_hashes,
&mut seq_hashes,
&mut active,
&mut scores,
current_pos + 1,
next_pos + 1,
false,
);
current_pos = next_pos;
}
}
let final_score = len as u32;
for worker in active {
scores.scores.insert(worker, final_score);
}
scores
}
fn binary_search_matches(
&self,
local_hashes: &[LocalBlockHash],
early_exit: bool,
) -> OverlapScores {
let mut scores = OverlapScores::new();
if local_hashes.is_empty() {
return scores;
}
let mut seq_hashes: Vec<ExternalSequenceBlockHash> = Vec::with_capacity(local_hashes.len());
let Some(mut active) =
self.get_workers_lazy(0, local_hashes[0], &mut seq_hashes, local_hashes)
else {
return scores;
};
if active.is_empty() {
return scores;
}
if early_exit {
for worker in &active {
scores.scores.insert(*worker, 1);
}
return scores;
}
let len = local_hashes.len();
let mut frontier = 0;
while frontier < len - 1 && !active.is_empty() {
if self
.count_workers_at(
len - 1,
local_hashes[len - 1],
&mut seq_hashes,
local_hashes,
)
.unwrap_or(0)
== active.len()
{
break;
}
let mut lo = frontier;
let mut hi = len - 1;
if hi - lo <= self.jump_size {
self.linear_scan_drain(
local_hashes,
&mut seq_hashes,
&mut active,
&mut scores,
lo + 1,
hi + 1,
false,
);
frontier = hi;
continue;
}
while hi - lo > 1 {
let mid = lo + (hi - lo) / 2;
if self
.count_workers_at(mid, local_hashes[mid], &mut seq_hashes, local_hashes)
.unwrap_or(0)
== active.len()
{
lo = mid;
} else {
hi = mid;
}
}
let drain_pos = hi;
Self::ensure_seq_hash_computed(&mut seq_hashes, drain_pos, local_hashes);
let drain_seq_hash = seq_hashes[drain_pos];
match self
.index
.get(&(drain_pos, local_hashes[drain_pos]))
.as_ref()
.and_then(|entry| entry.get(drain_seq_hash))
{
Some(workers) => {
reconcile_active_workers(&mut active, workers, |worker| {
scores.scores.insert(worker, drain_pos as u32);
});
}
None => {
for worker in active.drain() {
scores.scores.insert(worker, drain_pos as u32);
}
}
}
frontier = drain_pos;
}
let final_score = len as u32;
for worker in active {
scores.scores.insert(worker, final_score);
}
scores
}
}
#[cfg(test)]
mod tests {
use super::{LevelIndex, PositionalIndexer, SearchMode};
use crate::protocols::{LocalBlockHash, RouterEvent, WorkerWithDpRank};
use crate::test_utils::{assert_overlap_scores_eq, make_store_event};
use rustc_hash::FxHashMap;
fn local_hashes(hashes: &[u64]) -> Vec<LocalBlockHash> {
hashes.iter().copied().map(LocalBlockHash).collect()
}
fn populate(indexer: &PositionalIndexer, events: &[RouterEvent]) {
let mut worker_blocks: FxHashMap<WorkerWithDpRank, LevelIndex> = FxHashMap::default();
for ev in events {
indexer
.apply_event(&mut worker_blocks, ev.clone(), None)
.expect("apply_event should succeed");
}
}
#[test]
fn binary_matches_strided_both_early_exit_settings() {
let indexer = PositionalIndexer::new_with_mode(8, SearchMode::Strided);
populate(
&indexer,
&[
make_store_event(0, &[1, 2, 3, 4, 5, 6, 7, 8, 9, 10, 11, 12]),
make_store_event(1, &[1, 2, 3, 99, 100]),
make_store_event(2, &[1, 2, 3, 4, 5, 6, 7, 8, 200, 201]),
make_store_event(3, &[1, 2, 3, 4, 5, 6, 7, 8, 9, 10, 11, 12]),
],
);
let queries: &[&[u64]] = &[
&[],
&[1],
&[42], &[1, 2, 3],
&[1, 2, 3, 4, 5, 6, 7, 8],
&[1, 2, 3, 4, 5, 6, 7, 8, 9, 10, 11, 12],
&[1, 2, 3, 99, 100],
&[1, 2, 3, 4, 5, 6, 7, 8, 200, 201],
];
for q in queries {
let seq = local_hashes(q);
for early_exit in [false, true] {
let strided = indexer.jump_search_matches(&seq, early_exit);
let binary = indexer.binary_search_matches(&seq, early_exit);
assert_overlap_scores_eq(
&strided,
&binary,
&format!("q={q:?} early_exit={early_exit}"),
);
}
}
}
}