#![allow(clippy::all)]
use log::{debug, error, info, warn};
use std::sync::Arc;
use tokio::sync::RwLock;
use super::config::{BufferPoolConfig, CudaPoolConfig, ModelWeightPoolConfig};
use super::{BufferPool, BufferPoolStats, CudaMemoryPool, ModelPoolStats, ModelWeightPool};
use crate::error::VecboostError;
pub struct MemoryPoolManager {
buffer_pool: Arc<RwLock<BufferPool>>,
model_pool: Arc<RwLock<ModelWeightPool>>,
cuda_pool: Option<Arc<RwLock<CudaMemoryPool>>>,
config: MemoryPoolConfig,
}
#[derive(Debug, Clone)]
pub struct MemoryPoolConfig {
pub buffer_pool: BufferPoolConfig,
pub model_pool: ModelWeightPoolConfig,
pub cuda_pool: CudaPoolConfig,
}
impl MemoryPoolManager {
pub fn new(config: MemoryPoolConfig) -> Self {
info!("Creating MemoryPoolManager");
Self {
buffer_pool: Arc::new(RwLock::new(BufferPool::new(config.buffer_pool.clone()))),
model_pool: Arc::new(RwLock::new(ModelWeightPool::new(
"default".to_string(),
config.model_pool.clone(),
))),
cuda_pool: None, config,
}
}
pub async fn initialize_cuda_pool(&mut self, device_id: i32) -> Result<(), VecboostError> {
if !self.config.cuda_pool.enabled {
info!("CUDA pool disabled, skipping initialization");
return Ok(());
}
info!("Initializing CUDA pool for device {}...", device_id);
let pool = CudaMemoryPool::new(device_id, self.config.cuda_pool.clone()).map_err(|e| {
VecboostError::ConfigError(format!("Failed to initialize CUDA pool: {}", e))
})?;
self.cuda_pool = Some(Arc::new(RwLock::new(pool)));
info!("CUDA pool initialized successfully");
Ok(())
}
pub async fn initialize_all(
&mut self,
device: Option<candle_core::Device>,
cuda_device_id: Option<i32>,
) -> Result<(), VecboostError> {
info!("Initializing all memory pools...");
if self.config.cuda_pool.enabled && cuda_device_id.is_none() {
return Err(VecboostError::ConfigError(
"CUDA pool requires device_id but none provided".to_string(),
));
}
let mut initialized_pools = Vec::new();
if self.config.buffer_pool.enabled {
let mut buffer_pool = self.buffer_pool.write().await;
buffer_pool.preallocate();
initialized_pools.push("buffer_pool");
info!("Buffer pool initialized");
}
if self.config.model_pool.enabled {
initialized_pools.push("model_pool");
info!("Model weight pool initialized");
}
if let Some(device_id) = cuda_device_id {
match self.initialize_cuda_pool(device_id).await {
Ok(_) => {
initialized_pools.push("cuda_pool");
info!("CUDA pool initialized");
}
Err(e) => {
error!("Failed to initialize CUDA pool: {}, rolling back", e);
self.clear_all().await;
return Err(VecboostError::ConfigError(format!(
"Failed to initialize CUDA pool: {}. Rollback completed.",
e
)));
}
}
}
info!(
"All memory pools initialized successfully: {:?}",
initialized_pools
);
Ok(())
}
pub async fn get_buffer_pool(&self) -> Arc<RwLock<BufferPool>> {
Arc::clone(&self.buffer_pool)
}
pub async fn get_model_pool(&self) -> Arc<RwLock<ModelWeightPool>> {
Arc::clone(&self.model_pool)
}
pub async fn get_cuda_pool(&self) -> Option<Arc<RwLock<CudaMemoryPool>>> {
self.cuda_pool.clone()
}
pub async fn get_memory_stats(&self) -> MemoryPoolStats {
let buffer_pool = self.buffer_pool.read().await;
let model_pool = self.model_pool.read().await;
let mut stats = MemoryPoolStats {
buffer_pool_enabled: self.config.buffer_pool.enabled,
buffer_pool_stats: Some(buffer_pool.get_stats()),
model_pool_enabled: self.config.model_pool.enabled,
model_pool_stats: Some(model_pool.get_stats()),
cuda_pool_enabled: self.config.cuda_pool.enabled,
cuda_pool_stats: None,
};
if let Some(ref cuda_pool) = self.cuda_pool {
let pool = cuda_pool.read().await;
let (used, total) = pool.get_memory_usage();
stats.cuda_pool_stats = Some(CudaPoolStats {
used_memory_mb: used / 1024 / 1024,
total_memory_mb: total / 1024 / 1024,
memory_usage_percent: pool.get_memory_usage_percent(),
});
}
stats
}
pub async fn clear_all(&self) {
info!("Clearing all memory pools...");
{
let mut buffer_pool = self.buffer_pool.write().await;
buffer_pool.clear();
}
{
let mut model_pool = self.model_pool.write().await;
model_pool.clear();
}
if let Some(ref cuda_pool) = self.cuda_pool {
let mut pool = cuda_pool.write().await;
pool.clear();
}
info!("All memory pools cleared");
}
}
#[derive(Debug, Clone)]
pub struct MemoryPoolStats {
pub buffer_pool_enabled: bool,
pub buffer_pool_stats: Option<super::BufferPoolStats>,
pub model_pool_enabled: bool,
pub model_pool_stats: Option<super::ModelPoolStats>,
pub cuda_pool_enabled: bool,
pub cuda_pool_stats: Option<CudaPoolStats>,
}
#[derive(Debug, Clone)]
pub struct CudaPoolStats {
pub used_memory_mb: u64,
pub total_memory_mb: u64,
pub memory_usage_percent: f64,
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_pool_manager_creation() {
let config = MemoryPoolConfig {
buffer_pool: BufferPoolConfig::default(),
model_pool: ModelWeightPoolConfig::default(),
cuda_pool: CudaPoolConfig::default(),
};
let manager = MemoryPoolManager::new(config);
assert!(manager.cuda_pool.is_none());
}
#[tokio::test]
async fn test_get_buffer_pool() {
let config = MemoryPoolConfig {
buffer_pool: BufferPoolConfig::default(),
model_pool: ModelWeightPoolConfig::default(),
cuda_pool: CudaPoolConfig::default(),
};
let manager = MemoryPoolManager::new(config);
let buffer_pool = manager.get_buffer_pool().await;
let stats = buffer_pool.read().await.get_stats();
assert_eq!(stats.text_allocations, 0);
}
#[tokio::test]
async fn test_get_model_pool() {
let config = MemoryPoolConfig {
buffer_pool: BufferPoolConfig::default(),
model_pool: ModelWeightPoolConfig::default(),
cuda_pool: CudaPoolConfig::default(),
};
let manager = MemoryPoolManager::new(config);
let model_pool = manager.get_model_pool().await;
let stats = model_pool.read().await.get_stats();
assert_eq!(stats.loaded_models, 0);
}
#[tokio::test]
async fn test_initialize_all() {
let config = MemoryPoolConfig {
buffer_pool: BufferPoolConfig::default(),
model_pool: ModelWeightPoolConfig::default(),
cuda_pool: CudaPoolConfig {
enabled: false, ..Default::default()
},
};
let mut manager = MemoryPoolManager::new(config);
let result = manager.initialize_all(None, None).await;
assert!(result.is_ok(), "initialize_all failed: {:?}", result);
let stats = manager.get_memory_stats().await;
assert!(stats.buffer_pool_stats.is_some());
assert!(stats.model_pool_stats.is_some());
}
#[tokio::test]
async fn test_clear_all() {
let config = MemoryPoolConfig {
buffer_pool: BufferPoolConfig::default(),
model_pool: ModelWeightPoolConfig::default(),
cuda_pool: CudaPoolConfig::default(),
};
let manager = MemoryPoolManager::new(config);
manager.clear_all().await;
}
#[tokio::test]
async fn test_initialize_cuda_pool_disabled() {
let config = MemoryPoolConfig {
buffer_pool: BufferPoolConfig::default(),
model_pool: ModelWeightPoolConfig::default(),
cuda_pool: CudaPoolConfig {
enabled: false,
..Default::default()
},
};
let mut manager = MemoryPoolManager::new(config);
let result = manager.initialize_cuda_pool(0).await;
assert!(result.is_ok());
assert!(manager.cuda_pool.is_none());
}
#[tokio::test]
async fn test_initialize_cuda_pool_enabled() {
let config = MemoryPoolConfig {
buffer_pool: BufferPoolConfig::default(),
model_pool: ModelWeightPoolConfig::default(),
cuda_pool: CudaPoolConfig {
enabled: true,
max_memory_mb: 1024,
},
};
let mut manager = MemoryPoolManager::new(config);
let result = manager.initialize_cuda_pool(0).await;
assert!(result.is_ok());
assert!(manager.cuda_pool.is_some());
}
#[tokio::test]
async fn test_get_cuda_pool_returns_none_before_init() {
let config = MemoryPoolConfig {
buffer_pool: BufferPoolConfig::default(),
model_pool: ModelWeightPoolConfig::default(),
cuda_pool: CudaPoolConfig::default(),
};
let manager = MemoryPoolManager::new(config);
assert!(manager.get_cuda_pool().await.is_none());
}
#[tokio::test]
async fn test_initialize_all_error_cuda_enabled_no_device_id() {
let config = MemoryPoolConfig {
buffer_pool: BufferPoolConfig::default(),
model_pool: ModelWeightPoolConfig::default(),
cuda_pool: CudaPoolConfig {
enabled: true,
..Default::default()
},
};
let mut manager = MemoryPoolManager::new(config);
let result = manager.initialize_all(None, None).await;
assert!(result.is_err());
match result.unwrap_err() {
VecboostError::ConfigError(msg) => {
assert!(msg.contains("CUDA pool requires device_id"));
}
other => panic!("Expected ConfigError, got {:?}", other),
}
}
#[tokio::test]
async fn test_initialize_all_all_enabled() {
let config = MemoryPoolConfig {
buffer_pool: BufferPoolConfig {
enabled: true,
text_buffer_sizes: vec![16],
vector_buffer_sizes: vec![16],
pool_size_per_size: 2,
},
model_pool: ModelWeightPoolConfig {
enabled: true,
..Default::default()
},
cuda_pool: CudaPoolConfig {
enabled: true,
max_memory_mb: 1024,
},
};
let mut manager = MemoryPoolManager::new(config);
let result = manager
.initialize_all(Some(candle_core::Device::Cpu), Some(0))
.await;
assert!(result.is_ok(), "initialize_all failed: {:?}", result);
assert!(manager.cuda_pool.is_some());
}
#[tokio::test]
async fn test_get_memory_stats_with_all_pools() {
let config = MemoryPoolConfig {
buffer_pool: BufferPoolConfig {
enabled: true,
text_buffer_sizes: vec![16],
vector_buffer_sizes: vec![16],
pool_size_per_size: 2,
},
model_pool: ModelWeightPoolConfig {
enabled: true,
..Default::default()
},
cuda_pool: CudaPoolConfig {
enabled: true,
max_memory_mb: 1024,
},
};
let mut manager = MemoryPoolManager::new(config);
manager
.initialize_all(Some(candle_core::Device::Cpu), Some(0))
.await
.unwrap();
let stats = manager.get_memory_stats().await;
assert!(stats.buffer_pool_enabled);
assert!(stats.model_pool_enabled);
assert!(stats.cuda_pool_enabled);
assert!(stats.buffer_pool_stats.is_some());
assert!(stats.model_pool_stats.is_some());
assert!(stats.cuda_pool_stats.is_some());
let buffer_stats = stats.buffer_pool_stats.unwrap();
assert!(buffer_stats.current_text_pool_size > 0);
let cuda_stats = stats.cuda_pool_stats.unwrap();
assert_eq!(cuda_stats.used_memory_mb, 0);
#[cfg(feature = "cuda")]
assert_eq!(cuda_stats.total_memory_mb, 1024);
#[cfg(not(feature = "cuda"))]
assert_eq!(cuda_stats.total_memory_mb, 0);
assert!((cuda_stats.memory_usage_percent - 0.0).abs() < 0.001);
}
#[tokio::test]
async fn test_get_memory_stats_disabled_pools() {
let config = MemoryPoolConfig {
buffer_pool: BufferPoolConfig {
enabled: false,
..Default::default()
},
model_pool: ModelWeightPoolConfig {
enabled: false,
..Default::default()
},
cuda_pool: CudaPoolConfig {
enabled: false,
..Default::default()
},
};
let manager = MemoryPoolManager::new(config);
let stats = manager.get_memory_stats().await;
assert!(!stats.buffer_pool_enabled);
assert!(!stats.model_pool_enabled);
assert!(!stats.cuda_pool_enabled);
assert!(stats.cuda_pool_stats.is_none());
assert!(stats.buffer_pool_stats.is_some());
assert!(stats.model_pool_stats.is_some());
}
#[tokio::test]
async fn test_clear_all_after_initialization() {
let config = MemoryPoolConfig {
buffer_pool: BufferPoolConfig {
enabled: true,
text_buffer_sizes: vec![16],
vector_buffer_sizes: vec![16],
pool_size_per_size: 2,
},
model_pool: ModelWeightPoolConfig::default(),
cuda_pool: CudaPoolConfig {
enabled: true,
..Default::default()
},
};
let mut manager = MemoryPoolManager::new(config);
manager
.initialize_all(Some(candle_core::Device::Cpu), Some(0))
.await
.unwrap();
manager.clear_all().await;
let stats = manager.get_memory_stats().await;
let buffer_stats = stats.buffer_pool_stats.as_ref().unwrap();
assert_eq!(buffer_stats.current_text_pool_size, 0);
assert_eq!(buffer_stats.current_vector_pool_size, 0);
}
}