#![allow(unsafe_code)]
use crate::error::{LinuxError, Result, UmemError};
use std::sync::atomic::{AtomicU64, Ordering};
#[derive(Debug, Clone)]
pub struct UmemConfig {
pub size: usize,
pub hugepage: bool,
pub locked: bool,
pub shared: bool,
}
impl Default for UmemConfig {
fn default() -> Self {
Self {
size: 0,
hugepage: false,
locked: true,
shared: false,
}
}
}
#[derive(Debug, Clone)]
pub struct UmemRegion {
pub addr: *mut u8,
pub size: usize,
pub hugepage: bool,
pub locked: bool,
pub page_offset: u64,
}
unsafe impl Send for UmemRegion {}
unsafe impl Sync for UmemRegion {}
pub struct UmemManager {
region: Option<UmemRegion>,
config: UmemConfig,
ref_count: AtomicU64,
initialized: bool,
}
impl std::fmt::Debug for UmemManager {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("UmemManager")
.field("config", &self.config)
.field("initialized", &self.initialized)
.field("ref_count", &self.ref_count)
.finish()
}
}
impl UmemManager {
pub fn new(config: UmemConfig) -> Result<Self> {
if config.size == 0 {
return Err(LinuxError::Umem(UmemError::InsufficientSize {
actual: 0,
required: 4096,
}));
}
let page_size = crate::page_size();
if !config.size.is_multiple_of(page_size) {
return Err(LinuxError::Umem(UmemError::NotAligned {
actual: config.size,
expected: page_size,
}));
}
Ok(Self {
region: None,
config,
ref_count: AtomicU64::new(0),
initialized: false,
})
}
pub fn create(&mut self) -> Result<()> {
if self.initialized {
return Err(LinuxError::Umem(UmemError::AlreadyCreated));
}
let page_size = crate::page_size();
let mut flags = libc::MAP_PRIVATE | libc::MAP_ANONYMOUS;
if self.config.shared {
flags = libc::MAP_SHARED | libc::MAP_ANONYMOUS;
}
let region_size = if self.config.hugepage {
let hugepage_size = 2 * 1024 * 1024;
self.config
.size
.checked_add(hugepage_size - 1)
.map(|v| v & !(hugepage_size - 1))
.ok_or_else(|| {
LinuxError::InsufficientResources(
"UMEM size 对齐 HugePage 时溢出".to_string(),
)
})?
} else {
self.config.size
};
if self.config.hugepage {
flags |= libc::MAP_HUGETLB;
}
let addr = unsafe {
libc::mmap(
std::ptr::null_mut(),
region_size,
libc::PROT_READ | libc::PROT_WRITE,
flags,
-1,
0,
)
};
if addr == libc::MAP_FAILED {
let err = std::io::Error::last_os_error();
return Err(LinuxError::Umem(UmemError::MmapFailed(format!(
"mmap failed: {}",
err
))));
}
let addr = addr as *mut u8;
if self.config.locked {
let lock_result = unsafe { libc::mlock(addr as *const libc::c_void, region_size) };
if lock_result != 0 {
let err = std::io::Error::last_os_error();
unsafe {
libc::munmap(addr as *mut libc::c_void, region_size);
}
return Err(LinuxError::Umem(UmemError::LockFailed(format!(
"mlock failed: {}",
err
))));
}
}
if self.config.hugepage {
let madvise_ret = unsafe {
libc::madvise(addr as *mut libc::c_void, region_size, libc::MADV_HUGEPAGE)
};
if madvise_ret != 0 {
tracing::debug!("madvise(MADV_HUGEPAGE) 失败: {}", std::io::Error::last_os_error());
}
}
let page_offset = (addr as usize / page_size) as u64;
self.region = Some(UmemRegion {
addr,
size: region_size,
hugepage: self.config.hugepage,
locked: self.config.locked,
page_offset,
});
self.initialized = true;
self.ref_count.store(1, Ordering::SeqCst);
Ok(())
}
pub fn region(&self) -> Option<&UmemRegion> {
self.region.as_ref()
}
pub fn as_ptr(&self) -> *mut u8 {
self.region.as_ref().map_or(std::ptr::null_mut(), |r| r.addr)
}
#[inline]
pub fn slice(&self, offset: usize, len: usize) -> Option<&[u8]> {
let end = offset.checked_add(len)?;
if end > self.config.size {
return None;
}
let addr = self.as_ptr();
if addr.is_null() {
return None;
}
Some(unsafe { std::slice::from_raw_parts(addr.cast_const().add(offset), len) })
}
#[allow(clippy::mut_from_ref)]
#[inline]
pub fn slice_mut(&self, offset: usize, len: usize) -> Option<&mut [u8]> {
let end = offset.checked_add(len)?;
if end > self.config.size {
return None;
}
let addr = self.as_ptr();
if addr.is_null() {
return None;
}
Some(unsafe { std::slice::from_raw_parts_mut(addr.add(offset), len) })
}
pub fn page_offset(&self) -> u64 {
self.region.as_ref().map_or(0, |r| r.page_offset)
}
pub fn size(&self) -> usize {
self.config.size
}
pub fn is_initialized(&self) -> bool {
self.initialized
}
pub fn incref(&self) -> Result<u64> {
if !self.initialized {
return Err(LinuxError::Umem(UmemError::NotCreated));
}
loop {
let current = self.ref_count.load(Ordering::Acquire);
if current == 0 {
return Err(LinuxError::Umem(UmemError::MunmapFailed(
"incref: region already munmapped (ref_count == 0, use-after-free prevented)"
.to_string(),
)));
}
match self.ref_count.compare_exchange(
current,
current + 1,
Ordering::SeqCst,
Ordering::Acquire,
) {
Ok(_) => return Ok(current + 1),
Err(_) => continue,
}
}
}
pub fn decref(&self) -> u64 {
self.ref_count.fetch_sub(1, Ordering::SeqCst) - 1
}
}
impl Drop for UmemManager {
fn drop(&mut self) {
if let Some(region) = &self.region {
let cas_result = self.ref_count.compare_exchange(
1,
0,
Ordering::AcqRel,
Ordering::Acquire,
);
if cas_result.is_ok() {
unsafe {
libc::munmap(region.addr as *mut libc::c_void, region.size);
}
}
}
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_umem_config_default() {
let config = UmemConfig::default();
assert_eq!(config.size, 0);
assert!(!config.hugepage);
assert!(config.locked);
assert!(!config.shared);
}
#[test]
fn test_umem_manager_creation() {
let config = UmemConfig {
size: 0,
..Default::default()
};
let result = UmemManager::new(config);
assert!(result.is_err());
}
#[test]
fn test_umem_manager_create() {
let config = UmemConfig {
size: 4096 * 256, hugepage: false,
locked: false, shared: false,
};
let mut manager = UmemManager::new(config).unwrap();
assert!(!manager.is_initialized());
let result = manager.create();
assert!(result.is_ok());
assert!(manager.is_initialized());
assert!(!manager.as_ptr().is_null());
assert_eq!(manager.size(), 4096 * 256);
}
#[test]
fn test_umem_alignment_check() {
let config = UmemConfig {
size: 4096 + 1, ..Default::default()
};
let result = UmemManager::new(config);
assert!(result.is_err());
}
#[test]
fn test_umem_double_create() {
let config = UmemConfig {
size: 4096 * 64,
hugepage: false,
locked: false,
shared: false,
};
let mut manager = UmemManager::new(config).unwrap();
manager.create().unwrap();
let result = manager.create();
assert!(result.is_err());
}
#[test]
fn test_umem_size_zero_rejected() {
let config = UmemConfig {
size: 0,
..Default::default()
};
let result = UmemManager::new(config);
assert!(result.is_err());
}
#[test]
fn test_umem_unaligned_size_rejected() {
let page_size = crate::page_size();
let config = UmemConfig {
size: page_size + 1,
..Default::default()
};
let result = UmemManager::new(config);
assert!(result.is_err());
}
#[test]
fn test_umem_aligned_size_accepted() {
let page_size = unsafe { libc::sysconf(libc::_SC_PAGESIZE) } as usize;
for multiplier in [1, 2, 4, 8, 16, 32, 64, 128, 256] {
let config = UmemConfig {
size: page_size * multiplier,
hugepage: false,
locked: false,
shared: false,
};
let result = UmemManager::new(config);
assert!(
result.is_ok(),
"Size {} ({} * {}) should be valid",
page_size * multiplier,
page_size,
multiplier
);
}
}
#[test]
fn test_umem_region_access_after_create() {
let config = UmemConfig {
size: 4096 * 128,
hugepage: false,
locked: false,
shared: false,
};
let mut manager = UmemManager::new(config).unwrap();
assert!(manager.region().is_none());
assert!(manager.as_ptr().is_null());
assert_eq!(manager.page_offset(), 0);
manager.create().unwrap();
assert!(manager.region().is_some());
assert!(!manager.as_ptr().is_null());
let region = manager.region().unwrap();
assert!(!region.addr.is_null());
assert!(region.size >= 4096 * 128);
assert!(!region.hugepage);
assert!(!region.locked);
}
#[test]
fn test_umem_size_method_returns_config_size() {
let config = UmemConfig {
size: 4096 * 64,
hugepage: false,
locked: false,
shared: false,
};
let manager = UmemManager::new(config.clone()).unwrap();
assert_eq!(manager.size(), config.size);
}
#[test]
fn test_umem_ref_counting() {
let config = UmemConfig {
size: 4096 * 64,
hugepage: false,
locked: false,
shared: false,
};
let mut manager = UmemManager::new(config).unwrap();
manager.create().unwrap();
let ref1 = manager.incref().unwrap();
assert_eq!(ref1, 2);
let ref2 = manager.incref().unwrap();
assert_eq!(ref2, 3);
let dec1 = manager.decref();
assert_eq!(dec1, 2);
let dec2 = manager.decref();
assert_eq!(dec2, 1);
}
#[test]
fn test_umem_incref_before_create_fails() {
let config = UmemConfig {
size: 4096 * 16,
hugepage: false,
locked: false,
shared: false,
};
let manager = UmemManager::new(config).unwrap();
let result = manager.incref();
assert!(result.is_err());
assert!(matches!(
result.unwrap_err(),
LinuxError::Umem(UmemError::NotCreated)
));
}
#[test]
fn test_umem_hugepage_align_overflow_fails() {
let config = UmemConfig {
size: usize::MAX - 4095,
hugepage: true,
locked: false,
shared: false,
};
let mut manager = UmemManager::new(config).unwrap();
let result = manager.create();
assert!(result.is_err(), "HugePage 对齐溢出必须 Fail-Closed");
}
#[test]
fn test_umem_config_clone() {
let config = UmemConfig {
size: 4096 * 32,
hugepage: true,
locked: true,
shared: true,
};
let cloned = config.clone();
assert_eq!(cloned.size, config.size);
assert_eq!(cloned.hugepage, config.hugepage);
assert_eq!(cloned.locked, config.locked);
assert_eq!(cloned.shared, config.shared);
}
#[test]
fn test_umem_manager_debug_format() {
let config = UmemConfig {
size: 4096 * 16,
hugepage: false,
locked: false,
shared: false,
};
let manager = UmemManager::new(config).unwrap();
let debug = format!("{:?}", manager);
assert!(debug.contains("UmemManager"));
assert!(debug.contains("initialized"));
assert!(debug.contains("ref_count"));
}
#[test]
fn test_umem_initialized_state() {
let config = UmemConfig {
size: 4096 * 32,
hugepage: false,
locked: false,
shared: false,
};
let mut manager = UmemManager::new(config).unwrap();
assert!(!manager.is_initialized());
manager.create().unwrap();
assert!(manager.is_initialized());
}
#[test]
fn test_umem_large_allocation() {
let config = UmemConfig {
size: 4096 * 1024, hugepage: false,
locked: false,
shared: false,
};
let mut manager = match UmemManager::new(config) {
Ok(m) => m,
Err(_) => return, };
let result = manager.create();
if result.is_ok() {
assert!(manager.is_initialized());
assert!(!manager.as_ptr().is_null());
}
}
#[test]
fn test_umem_shared_config_flag() {
let config = UmemConfig {
size: 4096 * 32,
hugepage: false,
locked: false,
shared: true,
};
let mut manager = UmemManager::new(config).unwrap();
let result = manager.create();
assert!(result.is_ok());
assert!(manager.is_initialized());
}
#[test]
fn test_umem_region_send_sync() {
fn assert_send_sync<T: Send + Sync>() {}
assert_send_sync::<UmemRegion>();
}
}