use std::sync::atomic::{AtomicUsize, Ordering};
use std::sync::Arc;
use tensor_wasm_core::mem_pool::{DriverMemPool, MemPoolError};
use tensor_wasm_tenant::TenantContext;
use crate::unified::{UnifiedBacking, UnifiedError, UvmAdvice};
#[derive(Debug, Clone, Default)]
pub struct FreeLog {
frees: Arc<AtomicUsize>,
}
impl FreeLog {
pub fn new() -> Self {
Self::default()
}
pub fn frees(&self) -> usize {
self.frees.load(Ordering::Acquire)
}
fn record_free(&self) {
self.frees.fetch_add(1, Ordering::AcqRel);
}
}
#[derive(Debug)]
pub struct MockUnifiedBacking {
bytes: Vec<u8>,
free_log: FreeLog,
}
impl MockUnifiedBacking {
pub fn alloc(size: usize, free_log: FreeLog) -> Self {
Self {
bytes: vec![0u8; size],
free_log,
}
}
pub fn try_alloc(size: usize, free_log: FreeLog, fail: bool) -> Result<Self, UnifiedError> {
if fail {
return Err(UnifiedError::Allocation(
"mock-cuda: injected allocation failure".into(),
));
}
Ok(Self::alloc(size, free_log))
}
}
impl Drop for MockUnifiedBacking {
fn drop(&mut self) {
self.free_log.record_free();
}
}
impl UnifiedBacking for MockUnifiedBacking {
fn len(&self) -> usize {
self.bytes.len()
}
fn as_slice(&self) -> &[u8] {
&self.bytes
}
fn as_mut_slice(&mut self) -> &mut [u8] {
&mut self.bytes
}
fn apply_advice(&self, _hint: UvmAdvice) -> Result<(), UnifiedError> {
Ok(())
}
fn prefetch_to_device(&self, _device_ord: u32) -> Result<(), UnifiedError> {
Ok(())
}
fn prefetch_to_host(&self) -> Result<(), UnifiedError> {
Ok(())
}
}
pub fn mock_alloc_with_tenant_context(
size: usize,
tenant_ctx: Arc<TenantContext>,
free_log: FreeLog,
inject_alloc_failure: bool,
) -> Result<MockUnifiedBacking, tensor_wasm_core::error::TensorWasmError> {
if size == 0 {
return Err(UnifiedError::ZeroSize.into());
}
tenant_ctx.consume_gpu_bytes(size as u64)?;
match MockUnifiedBacking::try_alloc(size, free_log, inject_alloc_failure) {
Ok(backing) => Ok(backing),
Err(e) => {
tenant_ctx.release_gpu_bytes(size as u64);
Err(e.into())
}
}
}
#[derive(Debug)]
pub struct MockDriverMemPool {
cap_bytes: std::sync::atomic::AtomicU64,
free_log: FreeLog,
next_handle: AtomicUsize,
fail_alloc: std::sync::atomic::AtomicBool,
}
impl MockDriverMemPool {
pub fn new(free_log: FreeLog) -> Self {
Self {
cap_bytes: std::sync::atomic::AtomicU64::new(0),
free_log,
next_handle: AtomicUsize::new(1),
fail_alloc: std::sync::atomic::AtomicBool::new(false),
}
}
pub fn set_alloc_failure(&self, fail: bool) {
self.fail_alloc.store(fail, Ordering::Release);
}
pub fn allocate(&self, _size: usize) -> Result<usize, UnifiedError> {
if self.fail_alloc.load(Ordering::Acquire) {
return Err(UnifiedError::Cuda(
"mock-cuda: cuMemAllocFromPoolAsync -> CUDA_ERROR_OUT_OF_MEMORY".into(),
));
}
Ok(self.next_handle.fetch_add(1, Ordering::AcqRel))
}
pub fn deallocate(&self, _handle: usize) -> Result<(), UnifiedError> {
self.free_log.record_free();
Ok(())
}
}
impl DriverMemPool for MockDriverMemPool {
fn set_release_threshold(&self, bytes: u64) -> Result<(), MemPoolError> {
self.cap_bytes.store(bytes, Ordering::Relaxed);
Ok(())
}
fn release_threshold(&self) -> Option<u64> {
Some(self.cap_bytes.load(Ordering::Relaxed))
}
}
#[derive(Debug)]
pub struct MockTenantPoolBuffer {
pool: Arc<MockDriverMemPool>,
handle: usize,
size: usize,
}
impl MockTenantPoolBuffer {
pub fn new(pool: Arc<MockDriverMemPool>, size: usize) -> Result<Self, UnifiedError> {
let handle = pool.allocate(size)?;
Ok(Self { pool, handle, size })
}
pub fn len(&self) -> usize {
self.size
}
pub fn is_empty(&self) -> bool {
self.size == 0
}
}
impl Drop for MockTenantPoolBuffer {
fn drop(&mut self) {
if let Err(e) = self.pool.deallocate(self.handle) {
tracing::error!(
target: "tensor_wasm_mem::mock_cuda",
error = ?e,
"mock cuMemFreeAsync failed in MockTenantPoolBuffer::drop",
);
}
}
}
#[cfg(test)]
mod tests {
use super::*;
use tensor_wasm_core::types::TenantId;
use tensor_wasm_tenant::TenantContext;
fn capped_ctx(cap: u64) -> Arc<TenantContext> {
Arc::new(
TenantContext::builder(TenantId(1))
.with_gpu_memory_bytes_cap(cap)
.build(),
)
}
#[test]
fn mock_backing_records_free_once_on_drop() {
let log = FreeLog::new();
{
let mut b = MockUnifiedBacking::alloc(64, log.clone());
assert_eq!(b.len(), 64);
b.as_mut_slice().fill(7);
assert!(b.as_slice().iter().all(|&v| v == 7));
assert_eq!(log.frees(), 0, "no free before drop");
}
assert_eq!(log.frees(), 1, "exactly one free recorded on drop");
}
#[test]
fn rollback_restores_tenant_counter_on_alloc_failure() {
let ctx = capped_ctx(4096);
let log = FreeLog::new();
let err = mock_alloc_with_tenant_context(1024, ctx.clone(), log.clone(), true)
.expect_err("injected alloc failure must surface");
assert_eq!(
ctx.gpu_bytes_in_use(),
0,
"tenant GPU counter must be rolled back after alloc failure"
);
assert_eq!(log.frees(), 0);
assert!(matches!(
err,
tensor_wasm_core::error::TensorWasmError::Serialization(_)
));
}
#[test]
fn success_path_consumes_then_drop_does_not_double_free() {
let ctx = capped_ctx(4096);
let log = FreeLog::new();
let backing = mock_alloc_with_tenant_context(1024, ctx.clone(), log.clone(), false)
.expect("alloc should succeed");
assert_eq!(ctx.gpu_bytes_in_use(), 1024, "consume recorded");
drop(backing);
assert_eq!(log.frees(), 1, "backing freed exactly once");
}
#[test]
fn tenant_pool_frees_on_drop() {
let log = FreeLog::new();
let pool = Arc::new(MockDriverMemPool::new(log.clone()));
{
let buf = MockTenantPoolBuffer::new(pool.clone(), 256).expect("pool alloc");
assert_eq!(buf.len(), 256);
assert_eq!(log.frees(), 0, "no free before drop");
}
assert_eq!(log.frees(), 1, "tenant-pool buffer freed on drop");
}
#[test]
fn tenant_pool_alloc_failure_surfaces_and_records_no_free() {
let log = FreeLog::new();
let pool = Arc::new(MockDriverMemPool::new(log.clone()));
pool.set_alloc_failure(true);
let err = MockTenantPoolBuffer::new(pool.clone(), 256)
.expect_err("injected over-cap failure must surface");
assert!(matches!(err, UnifiedError::Cuda(_)));
assert_eq!(log.frees(), 0, "failed alloc records no free");
}
#[test]
fn mock_pool_is_a_driver_mem_pool() {
let log = FreeLog::new();
let pool: Arc<dyn DriverMemPool> = Arc::new(MockDriverMemPool::new(log));
pool.set_release_threshold(2048).unwrap();
assert_eq!(pool.release_threshold(), Some(2048));
}
}