use std::path::PathBuf;
use serde::{Deserialize, Serialize};
use validator::Validate;
#[derive(Debug, Clone, Copy, Default, Serialize, Deserialize, PartialEq, Eq)]
#[serde(rename_all = "snake_case")]
pub enum ParallelismMode {
#[default]
TensorParallel,
ReplicatedData,
}
#[derive(Debug, Clone, Serialize, Deserialize, Validate, Default)]
pub struct HostCacheConfig {
pub cache_size_gb: Option<f64>,
pub num_blocks: Option<usize>,
}
impl HostCacheConfig {
pub fn compute_num_blocks(&self, bytes_per_block: usize) -> Option<usize> {
if bytes_per_block == 0 {
return None;
}
self.num_blocks.or_else(|| {
self.cache_size_gb.map(|gb| {
((gb * 1_000_000_000.0) / bytes_per_block as f64) as usize
})
})
}
pub fn is_enabled(&self) -> bool {
self.num_blocks.is_some() || self.cache_size_gb.is_some()
}
}
#[derive(Debug, Clone, Serialize, Deserialize, Validate, Default)]
pub struct DiskCacheConfig {
pub cache_size_gb: Option<f64>,
pub num_blocks: Option<usize>,
#[serde(default)]
pub use_gds: bool,
pub storage_path: Option<PathBuf>,
}
impl DiskCacheConfig {
pub fn compute_num_blocks(&self, bytes_per_block: usize) -> Option<usize> {
if bytes_per_block == 0 {
return None;
}
self.num_blocks.or_else(|| {
self.cache_size_gb.map(|gb| {
((gb * 1_000_000_000.0) / bytes_per_block as f64) as usize
})
})
}
pub fn is_enabled(&self) -> bool {
self.num_blocks.is_some() || self.cache_size_gb.is_some()
}
}
#[derive(Debug, Clone, Default, Serialize, Deserialize, Validate)]
pub struct CacheConfig {
#[serde(default)]
#[validate(nested)]
pub host: HostCacheConfig,
#[validate(nested)]
pub disk: Option<DiskCacheConfig>,
#[serde(default)]
pub parallelism: ParallelismMode,
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_host_cache_default() {
let config = HostCacheConfig::default();
assert!(config.cache_size_gb.is_none());
assert!(config.num_blocks.is_none());
assert!(!config.is_enabled());
}
#[test]
fn test_host_cache_explicit_blocks() {
let config = HostCacheConfig {
num_blocks: Some(1000),
cache_size_gb: Some(10.0), };
let bytes_per_block = 1_000_000;
assert_eq!(config.compute_num_blocks(bytes_per_block), Some(1000));
assert!(config.is_enabled());
}
#[test]
fn test_host_cache_from_size_gb() {
let config = HostCacheConfig {
num_blocks: None,
cache_size_gb: Some(10.0), };
let bytes_per_block = 1_000_000;
assert_eq!(config.compute_num_blocks(bytes_per_block), Some(10_000));
assert!(config.is_enabled());
}
#[test]
fn test_disk_cache_default() {
let config = DiskCacheConfig::default();
assert!(config.cache_size_gb.is_none());
assert!(config.num_blocks.is_none());
assert!(!config.use_gds);
assert!(config.storage_path.is_none());
assert!(!config.is_enabled());
}
#[test]
fn test_disk_cache_with_gds() {
let config = DiskCacheConfig {
num_blocks: Some(5000),
cache_size_gb: None,
use_gds: true,
storage_path: Some(PathBuf::from("/mnt/nvme/kv_cache")),
};
assert!(config.use_gds);
assert_eq!(
config.storage_path,
Some(PathBuf::from("/mnt/nvme/kv_cache"))
);
assert!(config.is_enabled());
}
#[test]
fn test_parallelism_mode_default() {
let mode = ParallelismMode::default();
assert_eq!(mode, ParallelismMode::TensorParallel);
}
#[test]
fn test_parallelism_mode_serde() {
let tp = ParallelismMode::TensorParallel;
let json = serde_json::to_string(&tp).unwrap();
assert_eq!(json, "\"tensor_parallel\"");
let rd = ParallelismMode::ReplicatedData;
let json = serde_json::to_string(&rd).unwrap();
assert_eq!(json, "\"replicated_data\"");
let mode: ParallelismMode = serde_json::from_str("\"tensor_parallel\"").unwrap();
assert_eq!(mode, ParallelismMode::TensorParallel);
let mode: ParallelismMode = serde_json::from_str("\"replicated_data\"").unwrap();
assert_eq!(mode, ParallelismMode::ReplicatedData);
}
#[test]
fn test_cache_config_with_parallelism() {
let config = CacheConfig {
host: HostCacheConfig::default(),
disk: None,
parallelism: ParallelismMode::ReplicatedData,
};
assert_eq!(config.parallelism, ParallelismMode::ReplicatedData);
}
#[test]
fn test_cache_config_default_parallelism() {
let config = CacheConfig::default();
assert_eq!(config.parallelism, ParallelismMode::TensorParallel);
}
}