use std::collections::{HashMap, HashSet};
use std::future::Future;
use std::ops::Range;
use std::sync::LazyLock;
use std::time::Duration;
use dynamo_tokens::{SequenceHash, Token, compute_hash_v2, compute_next_sequence_hash};
use rustc_hash::FxHashMap;
use serde::{Deserialize, Serialize};
use xxhash_rust::xxh3;
const fn default_track_prefill_tokens() -> bool {
true
}
pub const KV_EVENT_SUBJECT: &str = "kv-events";
pub const XXH3_SEED: u64 = 1337;
pub fn compute_block_hash(data: &[u8]) -> LocalBlockHash {
LocalBlockHash(compute_hash_v2(data, XXH3_SEED))
}
#[derive(Debug, Clone, Copy, Default)]
pub struct BlockHashOptions<'a> {
pub block_mm_infos: Option<&'a [Option<BlockExtraInfo>]>,
pub lora_name: Option<&'a str>,
pub is_eagle: Option<bool>,
}
#[inline]
fn hash_block_no_mm(chunk: &[u32], seed: u64, scratch_bytes: &mut Vec<u8>) -> LocalBlockHash {
#[cfg(target_endian = "little")]
{
let _ = scratch_bytes;
let chunk_bytes = unsafe {
std::slice::from_raw_parts(chunk.as_ptr().cast::<u8>(), std::mem::size_of_val(chunk))
};
LocalBlockHash(xxh3::xxh3_64_with_seed(chunk_bytes, seed))
}
#[cfg(not(target_endian = "little"))]
{
scratch_bytes.clear();
for &token in chunk {
scratch_bytes.extend_from_slice(&token.to_le_bytes());
}
LocalBlockHash(xxh3::xxh3_64_with_seed(scratch_bytes, seed))
}
}
pub const MM_PAD_SHIFT_VALUE: u64 = 1_000_000;
pub const MM_PAD_HASH_MASK: u64 = (1 << 30) - 1;
pub fn pad_value_for_mm_hash(mm_hash: u64) -> u32 {
(MM_PAD_SHIFT_VALUE + (mm_hash & MM_PAD_HASH_MASK)) as u32
}
pub fn compute_block_hash_for_seq(
tokens: &[u32],
kv_block_size: u32,
options: BlockHashOptions<'_>,
) -> Vec<LocalBlockHash> {
if kv_block_size == 0 {
return Vec::new();
}
let seed = match options.lora_name.filter(|n| !n.is_empty()) {
Some(name) => XXH3_SEED.wrapping_add(xxh3::xxh3_64(name.as_bytes())),
None => XXH3_SEED,
};
let is_eagle_flag = options.is_eagle.unwrap_or(false);
let stride = kv_block_size as usize;
let window_size = if is_eagle_flag { stride + 1 } else { stride };
let estimated_blocks = if is_eagle_flag {
tokens.len().saturating_sub(1) / stride
} else {
tokens.len() / stride
};
let mut hashes = Vec::with_capacity(estimated_blocks);
let mut bytes = Vec::with_capacity(window_size * std::mem::size_of::<u32>());
let mut mm_hashes = Vec::new();
let mut block_idx = 0;
let mut start = 0;
while start + window_size <= tokens.len() {
let chunk = &tokens[start..start + window_size];
if let Some(mm_infos) = options.block_mm_infos
&& let Some(Some(block_mm_info)) = mm_infos.get(block_idx)
{
bytes.clear();
for &token in chunk {
bytes.extend_from_slice(&token.to_le_bytes());
}
mm_hashes.clear();
mm_hashes.extend(block_mm_info.mm_objects.iter().map(|obj| obj.mm_hash));
mm_hashes.sort_unstable();
for &mm_hash in &mm_hashes {
bytes.extend_from_slice(&mm_hash.to_le_bytes());
}
hashes.push(LocalBlockHash(xxh3::xxh3_64_with_seed(&bytes, seed)));
} else {
hashes.push(hash_block_no_mm(chunk, seed, &mut bytes));
}
start += stride;
block_idx += 1;
}
hashes
}
#[inline]
pub fn compute_next_seq_hash(
parent_seq_hash: SequenceHash,
current_block_hash: LocalBlockHash,
) -> SequenceHash {
compute_next_sequence_hash(parent_seq_hash, current_block_hash.0)
}
pub fn compute_seq_hash_for_block(block_hashes: &[LocalBlockHash]) -> Vec<SequenceHash> {
if block_hashes.is_empty() {
return Vec::new();
}
let mut sequence_hashes = Vec::with_capacity(block_hashes.len());
sequence_hashes.push(block_hashes[0].0);
for i in 1..block_hashes.len() {
let parent_seq_hash = sequence_hashes[i - 1];
sequence_hashes.push(compute_next_seq_hash(parent_seq_hash, block_hashes[i]));
}
sequence_hashes
}
pub trait WorkerConfigLike {
fn data_parallel_start_rank(&self) -> u32;
fn data_parallel_size(&self) -> u32;
fn max_num_batched_tokens(&self) -> Option<u64>;
fn total_kv_blocks(&self) -> Option<u64>;
fn taints(&self) -> &HashSet<String> {
&EMPTY_WORKER_TAINTS
}
fn stable_routing_id(&self) -> Option<&str> {
None
}
fn topology_domains(&self) -> Option<&HashMap<String, String>> {
None
}
fn kv_transfer_domain(&self) -> Option<&str> {
None
}
fn kv_transfer_enforcement(&self) -> Option<KvTransferEnforcement> {
None
}
fn kv_transfer_preferred_weight(&self) -> Option<f32> {
None
}
}
#[derive(Debug, Clone, Copy, Serialize, Deserialize, PartialEq, Eq)]
#[serde(rename_all = "snake_case")]
pub enum KvTransferEnforcement {
Required,
Preferred,
}
#[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Default)]
pub struct RoutingConstraints {
#[serde(default, skip_serializing_if = "HashSet::is_empty")]
pub required_taints: HashSet<String>,
#[serde(default, skip_serializing_if = "HashMap::is_empty")]
pub preferred_taints: HashMap<String, f32>,
}
impl RoutingConstraints {
pub fn is_empty(&self) -> bool {
self.required_taints.is_empty() && self.preferred_taints.is_empty()
}
pub fn has_hard_constraints(&self) -> bool {
!self.required_taints.is_empty()
}
pub fn is_compatible_with_worker_taints(&self, worker_taints: &HashSet<String>) -> bool {
if self.required_taints.is_empty() {
return true;
}
self.required_taints
.iter()
.all(|taint| worker_taints.contains(taint))
}
pub fn preferred_taint_matches(&self, worker_taints: &HashSet<String>) -> usize {
if self.preferred_taints.is_empty() {
return 0;
}
self.preferred_taints
.keys()
.filter(|taint| worker_taints.contains(*taint))
.count()
}
pub fn preferred_taint_multiplier(&self, worker_taints: &HashSet<String>) -> Option<f64> {
if self.preferred_taints.is_empty() {
return None;
}
let bias = self
.preferred_taints
.iter()
.filter(|(taint, _)| worker_taints.contains(*taint))
.map(|(_, weight)| f64::from(*weight))
.sum::<f64>()
.tanh();
Some((-bias).exp())
}
}
static EMPTY_WORKER_TAINTS: LazyLock<HashSet<String>> = LazyLock::new(HashSet::new);
pub trait RouterEventSink: Send + Sync {
fn publish_event(&self, event: &RouterEvent)
-> impl Future<Output = anyhow::Result<()>> + Send;
}
pub type WorkerId = u64;
pub type DpRank = u32;
#[derive(Debug, Clone, Copy, Serialize, Deserialize, PartialEq, Eq, Hash, PartialOrd, Ord)]
pub struct WorkerWithDpRank {
pub worker_id: WorkerId,
pub dp_rank: DpRank,
}
impl WorkerWithDpRank {
pub fn new(worker_id: WorkerId, dp_rank: DpRank) -> Self {
Self { worker_id, dp_rank }
}
pub fn from_worker_id(worker_id: WorkerId) -> Self {
Self {
worker_id,
dp_rank: 0,
}
}
}
#[derive(Debug, Clone, Copy, Default, Serialize, Deserialize, PartialEq, Eq, Hash)]
#[serde(rename_all = "snake_case")]
pub enum StorageTier {
#[default]
Device,
HostPinned,
Disk,
External,
}
impl StorageTier {
pub fn from_kv_medium(medium: &str) -> Option<Self> {
match medium {
"GPU" | "DEVICE" => Some(Self::Device),
"CPU" | "CPU_PINNED" | "CPU_TIER1" => Some(Self::HostPinned),
"CPU_TIER2" | "DISK" | "NVME" => Some(Self::Disk),
"EXTERNAL" | "NETWORK" | "REMOTE" | "SHARED" => Some(Self::External),
_ => None,
}
}
pub fn from_kv_medium_or_default(medium: Option<&str>) -> Self {
medium
.and_then(Self::from_kv_medium)
.unwrap_or(Self::Device)
}
pub fn to_kv_medium(self) -> Option<&'static str> {
match self {
Self::Device => None,
Self::HostPinned => Some("CPU_PINNED"),
Self::Disk => Some("DISK"),
Self::External => Some("EXTERNAL"),
}
}
pub fn is_gpu(self) -> bool {
matches!(self, Self::Device)
}
}
#[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Eq, Hash)]
pub enum PlacementOwner {
LocalWorker(WorkerWithDpRank),
Shared,
}
#[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Eq, Hash)]
pub struct Placement {
pub owner: PlacementOwner,
pub tier: StorageTier,
}
impl Placement {
pub fn local_worker(worker_id: WorkerId, dp_rank: DpRank, tier: StorageTier) -> Self {
Self {
owner: PlacementOwner::LocalWorker(WorkerWithDpRank::new(worker_id, dp_rank)),
tier,
}
}
pub fn local_gpu(worker_id: WorkerId, dp_rank: DpRank) -> Self {
Self::local_worker(worker_id, dp_rank, StorageTier::Device)
}
pub fn is_local_gpu(&self) -> bool {
matches!(self.owner, PlacementOwner::LocalWorker(_)) && self.tier.is_gpu()
}
}
#[derive(Debug, Clone, Serialize, Deserialize, PartialEq)]
pub struct PlacementEvent {
pub placement: Placement,
pub event: KvCacheEvent,
}
impl PlacementEvent {
pub fn new(placement: Placement, event: KvCacheEvent) -> Self {
Self { placement, event }
}
pub fn local_gpu(worker_id: WorkerId, event: KvCacheEvent) -> Self {
Self::new(Placement::local_gpu(worker_id, event.dp_rank), event)
}
pub fn into_router_event(self) -> Option<RouterEvent> {
let PlacementOwner::LocalWorker(worker) = self.placement.owner else {
return None;
};
Some(RouterEvent::with_storage_tier(
worker.worker_id,
self.event,
self.placement.tier,
))
}
}
#[derive(Debug, Clone, Serialize, Deserialize)]
#[serde(tag = "method", rename_all = "snake_case")]
pub enum RouterRequest {
#[serde(rename = "new")]
New {
tokens: Vec<Token>,
#[serde(default, skip_serializing_if = "Option::is_none")]
block_mm_infos: Option<Vec<Option<BlockExtraInfo>>>,
#[serde(default, skip_serializing_if = "RoutingConstraints::is_empty")]
routing_constraints: RoutingConstraints,
#[serde(default)]
priority_jump: f64,
#[serde(default, skip_serializing_if = "is_zero")]
strict_priority: u32,
#[serde(default, skip_serializing_if = "Option::is_none")]
lora_name: Option<String>,
},
PotentialLoads {
tokens: Vec<Token>,
#[serde(default, skip_serializing_if = "Option::is_none")]
block_mm_infos: Option<Vec<Option<BlockExtraInfo>>>,
#[serde(default, skip_serializing_if = "Option::is_none")]
lora_name: Option<String>,
},
MarkPrefill {
#[serde(default, skip_serializing_if = "Option::is_none")]
request_id: Option<String>,
},
MarkFree {
#[serde(default, skip_serializing_if = "Option::is_none")]
request_id: Option<String>,
},
}
impl Default for RouterRequest {
fn default() -> Self {
RouterRequest::New {
tokens: vec![],
block_mm_infos: None,
routing_constraints: RoutingConstraints::default(),
priority_jump: 0.0,
strict_priority: 0,
lora_name: None,
}
}
}
fn is_zero(value: &u32) -> bool {
*value == 0
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct PotentialLoad {
pub worker_id: WorkerId,
pub dp_rank: DpRank,
pub potential_prefill_tokens: usize,
pub potential_decode_blocks: usize,
#[serde(default)]
pub active_requests: usize,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
#[serde(tag = "method", rename_all = "snake_case")]
pub enum RouterResponse {
New {
worker_id: WorkerId,
#[serde(default)]
dp_rank: DpRank,
overlap_blocks: u32,
},
QueueRejected {
rejection: crate::scheduling::QueueRejection,
},
PrefillMarked {
success: bool,
},
FreeMarked {
success: bool,
},
PotentialLoads {
loads: Vec<PotentialLoad>,
#[serde(default)]
pending_count: usize,
#[serde(default)]
pending_isl_tokens: usize,
},
}
#[derive(Debug)]
pub struct WorkerSelectionResult {
pub worker: WorkerWithDpRank,
pub required_blocks: u64,
pub effective_overlap_blocks: f64,
pub cached_tokens: usize,
}
#[derive(Debug, Clone, Serialize, Deserialize, Default, PartialEq)]
pub struct ActiveLoad {
pub worker_id: WorkerId,
#[serde(default)]
pub dp_rank: DpRank,
pub active_decode_blocks: Option<u64>,
pub active_prefill_tokens: Option<u64>,
#[serde(default)]
pub kv_used_blocks: Option<u64>,
}
#[derive(Clone, Copy, Debug, PartialEq, Eq, Hash, Ord, PartialOrd)]
pub struct LocalBlockHash(pub u64);
#[derive(Clone, Copy, Debug, PartialEq, Eq, Hash, Ord, PartialOrd)]
pub struct ExternalSequenceBlockHash(pub u64);
impl From<u64> for ExternalSequenceBlockHash {
fn from(value: u64) -> Self {
Self(value)
}
}
impl From<i64> for ExternalSequenceBlockHash {
fn from(value: i64) -> Self {
Self(value as u64)
}
}
#[derive(Serialize, Deserialize, Debug, Clone)]
pub struct PrefillEvent {
pub request_id: String,
pub worker_id: WorkerId,
pub data: PrefillEventData,
pub router_id: u64,
}
#[derive(Serialize, Deserialize, Debug, Clone)]
pub enum PrefillEventData {
NewPrefill(usize),
UpdatePrefill(usize),
CompletePrefill,
}
#[derive(Serialize, Deserialize, Debug, Clone)]
pub struct ActiveSequenceEvent {
pub request_id: String,
pub worker: WorkerWithDpRank,
pub data: ActiveSequenceEventData,
pub router_id: u64,
#[serde(default)]
pub lora_name: Option<String>,
}
#[derive(Serialize, Deserialize, Debug, Clone, Copy, PartialEq, Eq)]
pub struct PrefillLoadHint {
pub initial_effective_prefill_tokens: usize,
pub expected_prefill_duration: Option<Duration>,
}
#[derive(Serialize, Deserialize, Debug, Clone)]
pub enum ActiveSequenceEventData {
AddRequest {
token_sequence: Option<Vec<SequenceHash>>,
#[serde(default = "default_track_prefill_tokens")]
track_prefill_tokens: bool,
expected_output_tokens: Option<u32>,
#[serde(default)]
prefill_load_hint: Option<PrefillLoadHint>,
},
Free,
MarkPrefillCompleted,
}
#[derive(Serialize, Deserialize, Debug, Clone)]
pub struct ActiveBlockEvent {
pub request_id: String,
pub data: ActiveBlockEventData,
}
#[derive(Serialize, Deserialize, Debug, Clone)]
pub enum ActiveBlockEventData {
NewBlock(Vec<SequenceHash>),
FreeBlock,
}
#[derive(Serialize, Deserialize, Debug, Clone)]
pub struct KvCacheEvents {
pub events: Vec<KvCacheEvent>,
pub shutdown: bool,
}
#[derive(Serialize, Deserialize, Debug, Clone, PartialEq)]
pub struct KvCacheEvent {
pub event_id: u64,
pub data: KvCacheEventData,
#[serde(default)]
pub dp_rank: DpRank,
}
#[derive(Serialize, Deserialize, Debug, Clone, PartialEq)]
#[serde(rename_all = "snake_case")]
pub enum KvCacheEventData {
Stored(KvCacheStoreData),
Removed(KvCacheRemoveData),
Cleared,
}
#[derive(Serialize, Deserialize, Debug, Clone, PartialEq)]
pub struct KvCacheStoreData {
pub parent_hash: Option<ExternalSequenceBlockHash>,
#[serde(default)]
pub start_position: Option<u32>,
pub blocks: Vec<KvCacheStoredBlockData>,
}
#[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Eq)]
pub struct BlockMmObjectInfo {
pub mm_hash: u64,
pub offsets: Vec<(usize, usize)>,
}
#[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Eq)]
pub struct BlockExtraInfo {
pub mm_objects: Vec<BlockMmObjectInfo>,
}
#[derive(Debug, Clone, Serialize, Deserialize, PartialEq)]
pub struct RequestMmObjectInfo {
pub mm_hash: u64,
pub offsets: Vec<(usize, usize)>,
}
#[derive(Debug, Clone, Serialize, Deserialize, PartialEq)]
pub struct RequestExtraInfo {
pub mm_objects: Vec<RequestMmObjectInfo>,
}
impl RequestExtraInfo {
pub fn to_block_level(
&self,
block_size: usize,
total_tokens: usize,
) -> Vec<Option<BlockExtraInfo>> {
let num_blocks = total_tokens.div_ceil(block_size);
let mut block_infos: Vec<Option<BlockExtraInfo>> = vec![None; num_blocks];
for req_mm_obj in &self.mm_objects {
for (req_start, req_end) in &req_mm_obj.offsets {
let start_block = req_start / block_size;
let end_block = (req_end.saturating_sub(1)) / block_size;
let upper_bound = end_block.min(num_blocks - 1) + 1;
for (block_idx, block_info_opt) in block_infos
.iter_mut()
.enumerate()
.take(upper_bound)
.skip(start_block)
{
let block_start_global = block_idx * block_size;
let block_end_global = ((block_idx + 1) * block_size).min(total_tokens);
let local_start = (*req_start).max(block_start_global) - block_start_global;
let local_end = (*req_end).min(block_end_global) - block_start_global;
if local_start < local_end {
let block_info = block_info_opt
.get_or_insert_with(|| BlockExtraInfo { mm_objects: vec![] });
if let Some(existing) = block_info
.mm_objects
.iter_mut()
.find(|obj| obj.mm_hash == req_mm_obj.mm_hash)
{
existing.offsets.push((local_start, local_end));
} else {
block_info.mm_objects.push(BlockMmObjectInfo {
mm_hash: req_mm_obj.mm_hash,
offsets: vec![(local_start, local_end)],
});
}
}
}
}
}
block_infos
}
}
#[derive(Serialize, Deserialize, Debug, Clone, PartialEq)]
pub struct KvCacheStoredBlockData {
pub block_hash: ExternalSequenceBlockHash,
pub tokens_hash: LocalBlockHash,
#[serde(default)]
pub mm_extra_info: Option<BlockExtraInfo>,
}
#[derive(Serialize, Deserialize, Debug, Clone, PartialEq)]
pub struct KvCacheRemoveData {
pub block_hashes: Vec<ExternalSequenceBlockHash>,
}
impl Serialize for LocalBlockHash {
fn serialize<S>(&self, serializer: S) -> Result<S::Ok, S::Error>
where
S: serde::Serializer,
{
serializer.serialize_u64(self.0)
}
}
impl<'de> Deserialize<'de> for LocalBlockHash {
fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
where
D: serde::Deserializer<'de>,
{
let value = u64::deserialize(deserializer)?;
Ok(LocalBlockHash(value))
}
}
impl Serialize for ExternalSequenceBlockHash {
fn serialize<S>(&self, serializer: S) -> Result<S::Ok, S::Error>
where
S: serde::Serializer,
{
serializer.serialize_u64(self.0)
}
}
impl<'de> Deserialize<'de> for ExternalSequenceBlockHash {
fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
where
D: serde::Deserializer<'de>,
{
let value = u64::deserialize(deserializer)?;
Ok(ExternalSequenceBlockHash(value))
}
}
#[derive(Debug, thiserror::Error)]
pub enum KvCacheEventError {
#[error("Failed to find parent block")]
ParentBlockNotFound,
#[error("Failed to find block")]
BlockNotFound,
#[error("Invalid block sequence")]
InvalidBlockSequence,
}
#[derive(Debug, Clone, Serialize, Deserialize, PartialEq)]
pub struct RouterEvent {
pub worker_id: WorkerId,
#[serde(default)]
pub storage_tier: StorageTier,
pub event: KvCacheEvent,
}
impl RouterEvent {
pub fn new(worker_id: WorkerId, event: KvCacheEvent) -> Self {
Self::with_storage_tier(worker_id, event, StorageTier::Device)
}
pub fn with_storage_tier(
worker_id: WorkerId,
event: KvCacheEvent,
storage_tier: StorageTier,
) -> Self {
Self {
worker_id,
storage_tier,
event,
}
}
}
#[derive(Debug, Clone, Default)]
pub struct SharedCacheHits {
pub ranges: Vec<Range<u32>>,
pub total_hits: u32,
}
impl SharedCacheHits {
pub fn from_ranges(ranges: Vec<Range<u32>>) -> Self {
let total_hits = ranges.iter().map(|r| r.end - r.start).sum();
Self { ranges, total_hits }
}
pub fn from_hits(hits: &[bool]) -> Self {
let mut ranges = Vec::new();
let mut i = 0;
while i < hits.len() {
if hits[i] {
let start = i as u32;
while i < hits.len() && hits[i] {
i += 1;
}
ranges.push(start..i as u32);
} else {
i += 1;
}
}
Self::from_ranges(ranges)
}
pub fn hits_beyond(&self, from_position: u32) -> u32 {
self.ranges
.iter()
.map(|r| {
if r.end <= from_position {
0
} else if r.start >= from_position {
r.end - r.start
} else {
r.end - from_position
}
})
.sum()
}
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct OverlapScores {
pub scores: FxHashMap<WorkerWithDpRank, u32>,
pub frequencies: Vec<usize>,
}
impl Default for OverlapScores {
fn default() -> Self {
Self::new()
}
}
impl OverlapScores {
pub fn new() -> Self {
Self {
scores: FxHashMap::default(),
frequencies: Vec::new(),
}
}
pub fn update_scores<'a, I>(&mut self, workers: I)
where
I: IntoIterator<Item = &'a WorkerWithDpRank>,
{
for worker in workers {
let score = self.scores.entry(*worker).or_insert(0);
*score += 1;
}
}
}
#[derive(Debug, Clone)]
pub struct TokensWithHashes {
tokens: Vec<u32>,
block_size: u32,
block_mm_infos: Option<Vec<Option<BlockExtraInfo>>>,
lora_name: Option<String>,
block_hashes: Option<Vec<LocalBlockHash>>,
seq_hashes: Option<Vec<SequenceHash>>,
is_eagle: Option<bool>,
}
impl TokensWithHashes {
pub fn new(tokens: Vec<u32>, block_size: u32) -> Self {
Self {
tokens,
block_size,
block_mm_infos: None,
lora_name: None,
block_hashes: None,
seq_hashes: None,
is_eagle: None,
}
}
pub fn with_mm_infos(mut self, infos: Vec<Option<BlockExtraInfo>>) -> Self {
self.block_mm_infos = Some(infos);
self.invalidate_hashes();
self
}
pub fn with_lora_name(mut self, name: String) -> Self {
self.lora_name = Some(name);
self.invalidate_hashes();
self
}
pub fn with_is_eagle(mut self, is_eagle: bool) -> Self {
self.set_is_eagle(is_eagle);
self
}
pub fn set_is_eagle(&mut self, is_eagle: bool) {
let is_eagle = Some(is_eagle);
if self.is_eagle == is_eagle {
return;
}
self.is_eagle = is_eagle;
self.invalidate_hashes();
}
fn invalidate_hashes(&mut self) {
self.block_hashes = None;
self.seq_hashes = None;
}
pub fn tokens(&self) -> &[u32] {
&self.tokens
}
pub fn len(&self) -> usize {
self.tokens.len()
}
pub fn is_empty(&self) -> bool {
self.tokens.is_empty()
}
pub fn block_size(&self) -> u32 {
self.block_size
}
pub fn block_mm_infos(&self) -> Option<&[Option<BlockExtraInfo>]> {
self.block_mm_infos.as_deref()
}
pub fn get_or_compute_block_hashes(&mut self) -> &[LocalBlockHash] {
if self.block_hashes.is_none() {
self.block_hashes = Some(compute_block_hash_for_seq(
&self.tokens,
self.block_size,
BlockHashOptions {
block_mm_infos: self.block_mm_infos.as_deref(),
lora_name: self.lora_name.as_deref(),
is_eagle: self.is_eagle,
},
));
}
self.block_hashes.as_ref().unwrap()
}
pub fn get_or_compute_seq_hashes(&mut self) -> &[SequenceHash] {
if self.seq_hashes.is_none() {
let block_hashes = self.get_or_compute_block_hashes();
self.seq_hashes = Some(compute_seq_hash_for_block(block_hashes));
}
self.seq_hashes.as_ref().unwrap()
}
pub fn block_hashes(&self) -> Option<&[LocalBlockHash]> {
self.block_hashes.as_deref()
}
pub fn seq_hashes(&self) -> Option<&[SequenceHash]> {
self.seq_hashes.as_deref()
}
}
#[cfg(test)]
mod tests {
use super::*;
use rstest::rstest;
use serde_json;
#[test]
fn pad_value_matches_sglang_protocol() {
assert_eq!(MM_PAD_SHIFT_VALUE, 1_000_000);
assert_eq!(MM_PAD_HASH_MASK, (1u64 << 30) - 1);
assert_eq!(pad_value_for_mm_hash(0), MM_PAD_SHIFT_VALUE as u32);
let fits = (1u64 << 30) - 1;
assert_eq!(
pad_value_for_mm_hash(fits),
(MM_PAD_SHIFT_VALUE + fits) as u32
);
let overflow = (1u64 << 30) | 0xCAFE;
assert_eq!(
pad_value_for_mm_hash(overflow),
(MM_PAD_SHIFT_VALUE + 0xCAFE) as u32,
"high bits above the 30-bit mask must be discarded"
);
}
#[test]
fn test_router_event_new() {
let worker_id = 0;
let kv_cache_event = KvCacheEvent {
event_id: 1,
data: KvCacheEventData::Stored(KvCacheStoreData {
parent_hash: None,
start_position: None,
blocks: vec![KvCacheStoredBlockData {
block_hash: ExternalSequenceBlockHash(0),
mm_extra_info: None,
tokens_hash: LocalBlockHash(13226331709069118873),
}],
}),
dp_rank: 0,
};
let router_event = RouterEvent::new(worker_id, kv_cache_event);
assert_eq!(router_event.worker_id, worker_id);
assert_eq!(router_event.event.event_id, 1);
if let KvCacheEventData::Stored(store_op) = &router_event.event.data {
assert_eq!(store_op.blocks.len(), 1);
assert_eq!(
store_op.blocks[0].tokens_hash,
compute_block_hash(b"test data")
);
assert_eq!(store_op.blocks[0].block_hash, ExternalSequenceBlockHash(0));
} else {
panic!("Expected KvCacheEventData::Stored");
}
}
#[rstest]
#[case(11)]
#[case(32)]
#[case(64)]
fn test_compute_block_hash_for_seq(#[case] kv_block_size: u32) {
let sequence = (0..kv_block_size).collect::<Vec<u32>>();
let hashes =
compute_block_hash_for_seq(&sequence, kv_block_size, BlockHashOptions::default());
assert_eq!(hashes.len(), 1);
let sequence = (0..(kv_block_size + 1)).collect::<Vec<u32>>();
let hashes =
compute_block_hash_for_seq(&sequence, kv_block_size, BlockHashOptions::default());
assert_eq!(hashes.len(), 1);
let sequence = (0..(2 * kv_block_size + 1)).collect::<Vec<u32>>();
let hashes =
compute_block_hash_for_seq(&sequence, kv_block_size, BlockHashOptions::default());
assert_eq!(hashes.len(), 2);
}
#[test]
fn test_compute_next_seq_hash_matches_rolling_hash() {
let block_hashes = [LocalBlockHash(11), LocalBlockHash(22), LocalBlockHash(33)];
let seq_hashes = compute_seq_hash_for_block(&block_hashes);
assert_eq!(
seq_hashes[1],
compute_next_seq_hash(seq_hashes[0], block_hashes[1])
);
assert_eq!(
seq_hashes[2],
compute_next_seq_hash(seq_hashes[1], block_hashes[2])
);
}
#[test]
fn test_lora_name_produces_different_hash() {
let tokens: Vec<u32> = (0..4).collect();
let base = compute_block_hash_for_seq(&tokens, 4, BlockHashOptions::default());
let lora_a = compute_block_hash_for_seq(
&tokens,
4,
BlockHashOptions {
lora_name: Some("adapter-a"),
..Default::default()
},
);
let lora_b = compute_block_hash_for_seq(
&tokens,
4,
BlockHashOptions {
lora_name: Some("adapter-b"),
..Default::default()
},
);
assert_ne!(base[0], lora_a[0]);
assert_ne!(base[0], lora_b[0]);
assert_ne!(lora_a[0], lora_b[0]);
}
#[test]
fn test_lora_name_empty_string_normalized_to_none() {
let tokens: Vec<u32> = (0..4).collect();
let base = compute_block_hash_for_seq(&tokens, 4, BlockHashOptions::default());
let empty = compute_block_hash_for_seq(
&tokens,
4,
BlockHashOptions {
lora_name: Some(""),
..Default::default()
},
);
assert_eq!(
base, empty,
"empty lora_name should be treated as base model"
);
}
#[test]
fn test_tokens_with_hashes_lora() {
let tokens: Vec<u32> = (0..8).collect();
let mut base = TokensWithHashes::new(tokens.clone(), 4);
let base_hashes = base.get_or_compute_block_hashes().to_vec();
let mut with_lora =
TokensWithHashes::new(tokens, 4).with_lora_name("my-adapter".to_string());
let lora_hashes = with_lora.get_or_compute_block_hashes().to_vec();
assert_eq!(base_hashes.len(), lora_hashes.len());
for (b, l) in base_hashes.iter().zip(lora_hashes.iter()) {
assert_ne!(b, l);
}
}
#[test]
fn test_tokens_with_hashes_lora_change_recomputes_cached_hashes() {
let tokens: Vec<u32> = (0..8).collect();
let mut with_hashes = TokensWithHashes::new(tokens.clone(), 4);
let base_sequence_hashes = with_hashes.get_or_compute_seq_hashes().to_vec();
let mut with_hashes = with_hashes.with_lora_name("my-adapter".to_string());
let actual_block_hashes = with_hashes.get_or_compute_block_hashes().to_vec();
let actual_sequence_hashes = with_hashes.get_or_compute_seq_hashes().to_vec();
let expected_block_hashes = compute_block_hash_for_seq(
&tokens,
4,
BlockHashOptions {
lora_name: Some("my-adapter"),
..Default::default()
},
);
let expected_sequence_hashes = compute_seq_hash_for_block(&expected_block_hashes);
assert_eq!(actual_block_hashes, expected_block_hashes);
assert_eq!(actual_sequence_hashes, expected_sequence_hashes);
assert_ne!(actual_sequence_hashes, base_sequence_hashes);
}
#[test]
fn test_tokens_with_hashes_mm_change_recomputes_cached_hashes() {
let tokens: Vec<u32> = (0..4).collect();
let mm_infos = vec![
Some(BlockExtraInfo {
mm_objects: vec![BlockMmObjectInfo {
mm_hash: 42,
offsets: vec![(0, 1)],
}],
}),
None,
];
let mut with_hashes = TokensWithHashes::new(tokens.clone(), 2);
let text_sequence_hashes = with_hashes.get_or_compute_seq_hashes().to_vec();
let mut with_hashes = with_hashes.with_mm_infos(mm_infos.clone());
let actual_block_hashes = with_hashes.get_or_compute_block_hashes().to_vec();
let actual_sequence_hashes = with_hashes.get_or_compute_seq_hashes().to_vec();
let expected_block_hashes = compute_block_hash_for_seq(
&tokens,
2,
BlockHashOptions {
block_mm_infos: Some(&mm_infos),
..Default::default()
},
);
let expected_sequence_hashes = compute_seq_hash_for_block(&expected_block_hashes);
assert_eq!(actual_block_hashes, expected_block_hashes);
assert_eq!(actual_sequence_hashes, expected_sequence_hashes);
assert_ne!(actual_sequence_hashes, text_sequence_hashes);
}
#[test]
fn test_compute_block_hash_for_seq_eagle_windows() {
let tokens: Vec<u32> = (0..6).collect();
let default_hashes = compute_block_hash_for_seq(&tokens, 2, BlockHashOptions::default());
let eagle_hashes = compute_block_hash_for_seq(
&tokens,
2,
BlockHashOptions {
is_eagle: Some(true),
..Default::default()
},
);
let expected_first = compute_block_hash_for_seq(
&[0, 1, 2],
2,
BlockHashOptions {
is_eagle: Some(true),
..Default::default()
},
);
let expected_second = compute_block_hash_for_seq(
&[2, 3, 4],
2,
BlockHashOptions {
is_eagle: Some(true),
..Default::default()
},
);
assert_eq!(default_hashes.len(), 3);
assert_eq!(eagle_hashes.len(), 2);
assert_eq!(eagle_hashes, vec![expected_first[0], expected_second[0]]);
assert_ne!(default_hashes[0], eagle_hashes[0]);
}
#[test]
fn test_tokens_with_hashes_set_is_eagle_invalidates_cache() {
let tokens: Vec<u32> = (0..6).collect();
let mut with_hashes = TokensWithHashes::new(tokens, 2);
let default_hashes = with_hashes.get_or_compute_block_hashes().to_vec();
with_hashes.set_is_eagle(true);
let eagle_hashes = with_hashes.get_or_compute_block_hashes().to_vec();
let expected_first = compute_block_hash_for_seq(
&[0, 1, 2],
2,
BlockHashOptions {
is_eagle: Some(true),
..Default::default()
},
);
let expected_second = compute_block_hash_for_seq(
&[2, 3, 4],
2,
BlockHashOptions {
is_eagle: Some(true),
..Default::default()
},
);
assert_eq!(default_hashes.len(), 3);
assert_eq!(eagle_hashes.len(), 2);
assert_eq!(eagle_hashes, vec![expected_first[0], expected_second[0]]);
assert_ne!(default_hashes[0], eagle_hashes[0]);
}
#[test]
fn test_local_block_hash_serialization() {
let hash = LocalBlockHash(12345);
let serialized = serde_json::to_string(&hash).unwrap();
assert_eq!(serialized, "12345");
let deserialized: LocalBlockHash = serde_json::from_str(&serialized).unwrap();
assert_eq!(deserialized, hash);
}
#[test]
fn test_external_sequence_block_hash_serialization() {
let hash = ExternalSequenceBlockHash(67890);
let serialized = serde_json::to_string(&hash).unwrap();
assert_eq!(serialized, "67890");
let deserialized: ExternalSequenceBlockHash = serde_json::from_str(&serialized).unwrap();
assert_eq!(deserialized, hash);
}
#[test]
fn test_router_request_mark_free_backwards_compatible_deserialization() {
let request: RouterRequest = serde_json::from_str(r#"{"method":"mark_free"}"#).unwrap();
assert!(matches!(
request,
RouterRequest::MarkFree { request_id: None }
));
}
#[test]
fn test_shared_cache_hits_from_hits() {
let hits = SharedCacheHits::from_hits(&[true, true, true, true]);
assert_eq!(hits.ranges, vec![0..4]);
assert_eq!(hits.total_hits, 4);
let hits = SharedCacheHits::from_hits(&[true, false, true, true, false, true]);
assert_eq!(hits.ranges, vec![0..1, 2..4, 5..6]);
assert_eq!(hits.total_hits, 4);
let hits = SharedCacheHits::from_hits(&[false, false, false]);
assert!(hits.ranges.is_empty());
assert_eq!(hits.total_hits, 0);
let hits = SharedCacheHits::from_hits(&[]);
assert!(hits.ranges.is_empty());
assert_eq!(hits.total_hits, 0);
}
#[test]
fn test_shared_cache_hits_beyond() {
#[allow(clippy::single_range_in_vec_init)]
let hits = SharedCacheHits::from_ranges(vec![0..4]);
assert_eq!(hits.hits_beyond(2), 2);
assert_eq!(hits.hits_beyond(0), 4);
assert_eq!(hits.hits_beyond(4), 0);
assert_eq!(hits.hits_beyond(10), 0);
}
#[test]
fn test_shared_cache_hits_beyond_sparse() {
let hits = SharedCacheHits::from_ranges(vec![1..3, 5..8]);
assert_eq!(hits.total_hits, 5);
assert_eq!(hits.hits_beyond(0), 5);
assert_eq!(hits.hits_beyond(2), 4);
assert_eq!(hits.hits_beyond(3), 3);
assert_eq!(hits.hits_beyond(6), 2);
assert_eq!(hits.hits_beyond(8), 0);
}
#[test]
fn test_kv_transfer_enforcement_serde() {
assert_eq!(
serde_json::to_string(&KvTransferEnforcement::Required).unwrap(),
r#""required""#
);
assert_eq!(
serde_json::from_str::<KvTransferEnforcement>(r#""preferred""#).unwrap(),
KvTransferEnforcement::Preferred
);
assert!(serde_json::from_str::<KvTransferEnforcement>(r#""fallback""#).is_err());
}
#[test]
fn test_worker_config_like_topology_domains_default() {
struct MinimalConfig;
impl WorkerConfigLike for MinimalConfig {
fn data_parallel_start_rank(&self) -> u32 {
0
}
fn data_parallel_size(&self) -> u32 {
1
}
fn max_num_batched_tokens(&self) -> Option<u64> {
None
}
fn total_kv_blocks(&self) -> Option<u64> {
None
}
}
let config = MinimalConfig;
assert!(
config.topology_domains().is_none(),
"Default topology_domains() should return None"
);
assert!(
config.kv_transfer_domain().is_none(),
"Default kv_transfer_domain() should return None"
);
assert!(
config.kv_transfer_enforcement().is_none(),
"Default kv_transfer_enforcement() should return None"
);
assert!(
config.kv_transfer_preferred_weight().is_none(),
"Default kv_transfer_preferred_weight() should return None"
);
}
#[test]
fn test_router_request_mark_free_serialization_with_request_id() {
let request = RouterRequest::MarkFree {
request_id: Some("req-123".to_string()),
};
let serialized = serde_json::to_string(&request).unwrap();
let deserialized: RouterRequest = serde_json::from_str(&serialized).unwrap();
assert_eq!(
serialized,
r#"{"method":"mark_free","request_id":"req-123"}"#
);
assert!(matches!(
deserialized,
RouterRequest::MarkFree {
request_id: Some(ref request_id)
} if request_id == "req-123"
));
}
#[test]
fn test_router_request_new_serialization_with_priority_jump() {
let request = RouterRequest::New {
tokens: vec![1, 2, 3],
block_mm_infos: None,
routing_constraints: RoutingConstraints::default(),
priority_jump: 5.0,
strict_priority: 0,
lora_name: None,
};
let serialized = serde_json::to_string(&request).unwrap();
let deserialized: RouterRequest = serde_json::from_str(&serialized).unwrap();
assert_eq!(
serialized,
r#"{"method":"new","tokens":[1,2,3],"priority_jump":5.0}"#
);
assert!(matches!(
deserialized,
RouterRequest::New {
priority_jump,
..
} if priority_jump == 5.0
));
}
#[test]
fn test_router_request_new_serialization_with_lora_name() {
let request = RouterRequest::New {
tokens: vec![1, 2, 3],
block_mm_infos: None,
routing_constraints: RoutingConstraints::default(),
priority_jump: 0.0,
strict_priority: 0,
lora_name: Some("adapter-a".to_string()),
};
let serialized = serde_json::to_string(&request).unwrap();
let deserialized: RouterRequest = serde_json::from_str(&serialized).unwrap();
assert_eq!(
serialized,
r#"{"method":"new","tokens":[1,2,3],"priority_jump":0.0,"lora_name":"adapter-a"}"#
);
assert!(matches!(
deserialized,
RouterRequest::New {
tokens,
lora_name: Some(ref lora_name),
..
} if tokens == vec![1, 2, 3] && lora_name == "adapter-a"
));
}
#[test]
fn test_router_request_new_defaults_lora_name() {
let deserialized: RouterRequest =
serde_json::from_str(r#"{"method":"new","tokens":[1,2,3]}"#).unwrap();
assert!(matches!(
deserialized,
RouterRequest::New {
tokens,
lora_name: None,
..
} if tokens == vec![1, 2, 3]
));
}
#[test]
fn test_router_request_new_strict_priority_compatibility() {
let request = RouterRequest::New {
tokens: vec![1, 2, 3],
block_mm_infos: None,
routing_constraints: RoutingConstraints::default(),
priority_jump: 0.0,
strict_priority: 4,
lora_name: None,
};
let serialized = serde_json::to_string(&request).unwrap();
assert_eq!(
serialized,
r#"{"method":"new","tokens":[1,2,3],"priority_jump":0.0,"strict_priority":4}"#
);
let missing: RouterRequest =
serde_json::from_str(r#"{"method":"new","tokens":[1,2,3]}"#).unwrap();
assert!(matches!(
missing,
RouterRequest::New {
strict_priority: 0,
..
}
));
let zero = RouterRequest::New {
tokens: vec![1, 2, 3],
block_mm_infos: None,
routing_constraints: RoutingConstraints::default(),
priority_jump: 0.0,
strict_priority: 0,
lora_name: None,
};
assert_eq!(
serde_json::to_string(&zero).unwrap(),
r#"{"method":"new","tokens":[1,2,3],"priority_jump":0.0}"#
);
}
#[test]
fn test_router_request_potential_loads_serialization_with_lora_name() {
let request = RouterRequest::PotentialLoads {
tokens: vec![1, 2, 3],
block_mm_infos: None,
lora_name: Some("adapter-a".to_string()),
};
let serialized = serde_json::to_string(&request).unwrap();
let deserialized: RouterRequest = serde_json::from_str(&serialized).unwrap();
assert_eq!(
serialized,
r#"{"method":"potential_loads","tokens":[1,2,3],"lora_name":"adapter-a"}"#
);
assert!(matches!(
deserialized,
RouterRequest::PotentialLoads {
tokens,
block_mm_infos: None,
lora_name: Some(ref lora_name),
} if tokens == vec![1, 2, 3] && lora_name == "adapter-a"
));
}
#[test]
fn test_router_request_potential_loads_defaults_lora_name() {
let deserialized: RouterRequest =
serde_json::from_str(r#"{"method":"potential_loads","tokens":[1,2,3]}"#).unwrap();
assert!(matches!(
deserialized,
RouterRequest::PotentialLoads {
tokens,
block_mm_infos: None,
lora_name: None,
} if tokens == vec![1, 2, 3]
));
}
#[test]
fn test_router_request_mark_prefill_serialization_with_request_id() {
let request = RouterRequest::MarkPrefill {
request_id: Some("req-123".to_string()),
};
let serialized = serde_json::to_string(&request).unwrap();
let deserialized: RouterRequest = serde_json::from_str(&serialized).unwrap();
assert_eq!(
serialized,
r#"{"method":"mark_prefill","request_id":"req-123"}"#
);
assert!(matches!(
deserialized,
RouterRequest::MarkPrefill {
request_id: Some(ref request_id)
} if request_id == "req-123"
));
}
#[test]
fn test_potential_load_defaults_active_requests() {
let load = serde_json::from_str::<PotentialLoad>(
r#"{"worker_id":1,"dp_rank":0,"potential_prefill_tokens":16,"potential_decode_blocks":4}"#,
)
.unwrap();
assert_eq!(load.worker_id, 1);
assert_eq!(load.dp_rank, 0);
assert_eq!(load.potential_prefill_tokens, 16);
assert_eq!(load.potential_decode_blocks, 4);
assert_eq!(load.active_requests, 0);
}
#[test]
fn test_potential_load_serializes_active_requests() {
let load = PotentialLoad {
worker_id: 1,
dp_rank: 0,
potential_prefill_tokens: 16,
potential_decode_blocks: 4,
active_requests: 2,
};
assert_eq!(
serde_json::to_string(&load).unwrap(),
r#"{"worker_id":1,"dp_rank":0,"potential_prefill_tokens":16,"potential_decode_blocks":4,"active_requests":2}"#
);
}
}