use std::collections::HashMap;
use std::sync::Arc;
use crate::common::protocols::MockEngineArgs;
use dynamo_kv_router::config::KvRouterConfig;
use dynamo_kv_router::protocols::{
ActiveSequenceEvent, WorkerConfigLike, WorkerId, WorkerWithDpRank,
};
use dynamo_kv_router::scheduling::queue::DEFAULT_MAX_BATCHED_TOKENS;
use dynamo_kv_router::sequences::SchedulerLoadSnapshot;
use dynamo_kv_router::{
ActiveSequencesMultiWorker, DefaultWorkerSelector, LocalScheduler, SequencePublisher,
};
#[derive(Clone, Copy, Debug, Default)]
pub(super) struct ReplayNoopPublisher;
impl SequencePublisher for ReplayNoopPublisher {
fn enqueue_event(&self, _event: ActiveSequenceEvent) -> anyhow::Result<()> {
Ok(())
}
fn publish_scheduler_load(&self, _load: SchedulerLoadSnapshot) {}
fn observe_load(&self, _: &WorkerWithDpRank, _: &str, _: usize, _: usize) {}
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub(super) struct ReplayWorkerConfig {
pub(super) max_num_batched_tokens: u64,
pub(super) total_kv_blocks: u64,
pub(super) data_parallel_start_rank: u32,
pub(super) data_parallel_size: u32,
}
impl WorkerConfigLike for ReplayWorkerConfig {
fn data_parallel_start_rank(&self) -> u32 {
self.data_parallel_start_rank
}
fn data_parallel_size(&self) -> u32 {
self.data_parallel_size
}
fn max_num_batched_tokens(&self) -> Option<u64> {
Some(self.max_num_batched_tokens)
}
fn total_kv_blocks(&self) -> Option<u64> {
Some(self.total_kv_blocks)
}
}
pub(super) type ReplayScheduler =
LocalScheduler<ReplayNoopPublisher, ReplayWorkerConfig, DefaultWorkerSelector>;
pub(in crate::replay) fn replay_worker_config(args: &MockEngineArgs) -> ReplayWorkerConfig {
ReplayWorkerConfig {
max_num_batched_tokens: args
.max_num_batched_tokens
.map(|tokens| tokens as u64)
.unwrap_or(DEFAULT_MAX_BATCHED_TOKENS),
total_kv_blocks: args.num_gpu_blocks as u64,
data_parallel_start_rank: 0,
data_parallel_size: args.dp_size.max(1),
}
}
pub(super) fn replay_workers_with_configs(
args: &MockEngineArgs,
num_workers: usize,
) -> HashMap<WorkerId, ReplayWorkerConfig> {
let worker_config = replay_worker_config(args);
(0..num_workers)
.map(|worker_idx| (worker_idx as WorkerId, worker_config.clone()))
.collect()
}
pub(super) fn replay_slots(
args: &MockEngineArgs,
workers_with_configs: &HashMap<WorkerId, ReplayWorkerConfig>,
) -> Arc<ActiveSequencesMultiWorker<ReplayNoopPublisher>> {
let dp_range = workers_with_configs
.iter()
.map(|(&worker_id, config)| {
(
worker_id,
(config.data_parallel_start_rank, config.data_parallel_size),
)
})
.collect();
Arc::new(ActiveSequencesMultiWorker::new_without_expiry(
ReplayNoopPublisher,
args.block_size,
dp_range,
false,
0,
"replay",
))
}
pub(super) fn replay_selector(config: &KvRouterConfig) -> anyhow::Result<DefaultWorkerSelector> {
replay_selector_with_seed(config, None)
}
pub(super) fn replay_selector_with_seed(
config: &KvRouterConfig,
selector_seed: Option<u64>,
) -> anyhow::Result<DefaultWorkerSelector> {
if let Some(instance) = config
.selected_worker_selection_policy_instance()
.map_err(anyhow::Error::from)?
{
anyhow::bail!("custom worker-selection policy {instance:?} is not supported by replay");
}
Ok(match selector_seed {
#[cfg(feature = "replay-bench")]
Some(seed) => DefaultWorkerSelector::new_seeded(Some(config.clone()), "replay", seed),
#[cfg(not(feature = "replay-bench"))]
Some(_) => unreachable!("canonical KV Router replay requires the replay-bench feature"),
None => DefaultWorkerSelector::new(Some(config.clone()), "replay"),
})
}
pub(crate) fn replay_router_config(
args: &MockEngineArgs,
router_config: Option<KvRouterConfig>,
) -> KvRouterConfig {
let mut config = router_config.unwrap_or_default();
if let Some(policy) = args.router_queue_policy {
config.router_queue_policy = policy;
}
config
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn replay_selector_rejects_custom_worker_selection() {
let policy_file = tempfile::NamedTempFile::new().unwrap();
std::fs::write(
policy_file.path(),
r#"
worker_selection:
aggregated: custom
instances:
- name: custom
type: test
parameters: {}
"#,
)
.unwrap();
let config = KvRouterConfig {
router_policy_config: Some(policy_file.path().display().to_string()),
..Default::default()
};
let error = match replay_selector(&config) {
Err(error) => error,
Ok(_) => panic!("replay must reject a custom worker selector it cannot execute"),
};
assert!(
error
.to_string()
.contains("custom worker-selection policy \"custom\" is not supported by replay")
);
}
}