use std::collections::HashMap;
use std::collections::hash_map::Entry;
use dynamo_kv_router::protocols::{
ExternalSequenceBlockHash, KvCacheRemoveData, KvCacheStoreData, ResidencyDomain, StorageTier,
};
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub(super) enum EventDedupPolicy {
RefCounted,
SetLike,
}
pub(super) struct EventDedupFilter {
per_rank_tier:
HashMap<(u32, StorageTier, ResidencyDomain), HashMap<ExternalSequenceBlockHash, usize>>,
}
impl EventDedupFilter {
pub(super) fn new() -> Self {
Self {
per_rank_tier: HashMap::new(),
}
}
pub(super) fn track_store_in_domain(
&mut self,
dp_rank: u32,
storage_tier: StorageTier,
residency_domain: ResidencyDomain,
policy: EventDedupPolicy,
data: &KvCacheStoreData,
) {
if policy == EventDedupPolicy::SetLike {
return;
}
let refcounts = self
.per_rank_tier
.entry((dp_rank, storage_tier, residency_domain))
.or_default();
for block in &data.blocks {
*refcounts.entry(block.block_hash).or_insert(0) += 1;
}
}
pub(super) fn filter_remove_in_domain(
&mut self,
dp_rank: u32,
storage_tier: StorageTier,
residency_domain: ResidencyDomain,
policy: EventDedupPolicy,
mut data: KvCacheRemoveData,
) -> Option<KvCacheRemoveData> {
if policy == EventDedupPolicy::SetLike {
return (!data.block_hashes.is_empty()).then_some(data);
}
let refcounts = self
.per_rank_tier
.entry((dp_rank, storage_tier, residency_domain))
.or_default();
data.block_hashes.retain(|hash| {
match refcounts.entry(*hash) {
Entry::Occupied(mut entry) => {
*entry.get_mut() -= 1;
if *entry.get() == 0 {
entry.remove();
true } else {
false }
}
Entry::Vacant(_) => {
true }
}
});
if data.block_hashes.is_empty() {
None
} else {
Some(data)
}
}
pub(super) fn clear_rank_domain(
&mut self,
dp_rank: u32,
domain: ResidencyDomain,
policy: EventDedupPolicy,
) {
if policy == EventDedupPolicy::SetLike {
return;
}
self.per_rank_tier
.retain(|(tracked_dp_rank, _, tracked_domain), _| {
*tracked_dp_rank != dp_rank || *tracked_domain != domain
});
}
#[cfg(test)]
pub(super) fn track_store(
&mut self,
dp_rank: u32,
storage_tier: StorageTier,
data: &KvCacheStoreData,
) {
self.track_store_in_domain(
dp_rank,
storage_tier,
ResidencyDomain::Worker,
EventDedupPolicy::RefCounted,
data,
);
}
#[cfg(test)]
pub(super) fn filter_remove(
&mut self,
dp_rank: u32,
storage_tier: StorageTier,
data: KvCacheRemoveData,
) -> Option<KvCacheRemoveData> {
self.filter_remove_in_domain(
dp_rank,
storage_tier,
ResidencyDomain::Worker,
EventDedupPolicy::RefCounted,
data,
)
}
#[cfg(test)]
pub(super) fn clear_rank(&mut self, dp_rank: u32) {
self.per_rank_tier
.retain(|(tracked_dp_rank, _, _), _| *tracked_dp_rank != dp_rank);
}
}