use std::sync::Arc;
use zenoh::Wait;
use zenoh::shm::{BlockOn, GarbageCollect, PosixShmProviderBackend, ShmProvider, ZShmMut};
use zenoh_buffers::ZBuf;
pub const DEFAULT_SHM_POOL_SIZE: usize = 10 * 1024 * 1024;
pub const DEFAULT_SHM_THRESHOLD: usize = 512;
#[derive(Clone)]
pub struct ShmConfig {
pub(crate) provider: Arc<ShmProvider<PosixShmProviderBackend>>,
pub(crate) threshold: usize,
}
impl ShmConfig {
pub fn new(provider: Arc<ShmProvider<PosixShmProviderBackend>>) -> Self {
Self {
provider,
threshold: DEFAULT_SHM_THRESHOLD,
}
}
pub fn with_threshold(mut self, threshold: usize) -> Self {
self.threshold = threshold;
self
}
pub fn threshold(&self) -> usize {
self.threshold
}
pub fn provider(&self) -> &ShmProvider<PosixShmProviderBackend> {
&self.provider
}
pub fn from_env() -> zenoh::Result<Option<Self>> {
let has_pool_size = std::env::var("ZENOH_SHM_ALLOC_SIZE").is_ok();
let has_threshold = std::env::var("ZENOH_SHM_MESSAGE_SIZE_THRESHOLD").is_ok();
if !has_pool_size && !has_threshold {
return Ok(None);
}
let pool_size = std::env::var("ZENOH_SHM_ALLOC_SIZE")
.ok()
.and_then(|s| s.parse().ok())
.unwrap_or(DEFAULT_SHM_POOL_SIZE);
let threshold = std::env::var("ZENOH_SHM_MESSAGE_SIZE_THRESHOLD")
.ok()
.and_then(|s| s.parse().ok())
.unwrap_or(DEFAULT_SHM_THRESHOLD);
let provider = Arc::new(ShmProviderBuilder::new(pool_size).build()?);
Ok(Some(Self {
provider,
threshold,
}))
}
}
impl std::fmt::Debug for ShmConfig {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("ShmConfig")
.field("threshold", &self.threshold)
.field("provider", &"<ShmProvider>")
.finish()
}
}
pub struct ShmProviderBuilder {
size: usize,
}
impl ShmProviderBuilder {
pub fn new(size: usize) -> Self {
Self { size }
}
pub fn build(self) -> zenoh::Result<ShmProvider<PosixShmProviderBackend>> {
use zenoh::shm::ShmProviderBuilder as ZenohShmProviderBuilder;
ZenohShmProviderBuilder::default_backend(self.size)
.wait()
.map_err(|e| zenoh::Error::from(format!("Failed to create ShmProvider: {}", e)))
}
}
pub struct ShmWriter {
buffer: ZShmMut,
position: usize,
}
impl ShmWriter {
pub fn new(
provider: &ShmProvider<PosixShmProviderBackend>,
capacity: usize,
) -> zenoh::Result<Self> {
let buffer = provider
.alloc(capacity)
.with_policy::<BlockOn<GarbageCollect>>()
.wait()
.map_err(|e| zenoh::Error::from(format!("SHM allocation failed: {}", e)))?;
Ok(Self {
buffer,
position: 0,
})
}
#[inline]
pub fn position(&self) -> usize {
self.position
}
pub fn into_zbuf(self) -> zenoh::Result<ZBuf> {
Ok(ZBuf::from(self.buffer))
}
#[inline]
fn write_bytes(&mut self, bytes: &[u8]) {
let end = self.position + bytes.len();
assert!(
end <= self.buffer.len(),
"SHM buffer overflow: tried to write {} bytes at position {} but buffer size is {}",
bytes.len(),
self.position,
self.buffer.len()
);
self.buffer[self.position..end].copy_from_slice(bytes);
self.position = end;
}
}
impl hiroz_cdr::CdrBuffer for ShmWriter {
#[inline(always)]
fn extend_from_slice(&mut self, data: &[u8]) {
self.write_bytes(data);
}
#[inline(always)]
fn push(&mut self, byte: u8) {
self.write_bytes(&[byte]);
}
#[inline(always)]
fn len(&self) -> usize {
self.position
}
#[inline(always)]
fn reserve(&mut self, _additional: usize) {
}
#[inline(always)]
fn clear(&mut self) {
self.position = 0;
}
fn append_zbuf(&mut self, zbuf: &ZBuf) {
use zenoh_buffers::buffer::SplitBuffer;
let bytes = zbuf.contiguous();
self.write_bytes(&bytes);
}
}
#[cfg(test)]
mod tests {
use super::*;
use serial_test::serial;
#[test]
fn test_shm_config_creation() {
let provider = Arc::new(
ShmProviderBuilder::new(1024 * 1024)
.build()
.expect("Failed to create SHM provider"),
);
let config = ShmConfig::new(provider);
assert_eq!(config.threshold(), DEFAULT_SHM_THRESHOLD);
}
#[test]
fn test_shm_config_with_threshold() {
let provider = Arc::new(
ShmProviderBuilder::new(1024 * 1024)
.build()
.expect("Failed to create SHM provider"),
);
let config = ShmConfig::new(provider).with_threshold(10_000);
assert_eq!(config.threshold(), 10_000);
}
#[test]
fn test_shm_provider_builder() {
let provider = ShmProviderBuilder::new(2 * 1024 * 1024)
.build()
.expect("Failed to create SHM provider");
let buf = provider.alloc(1024).wait();
assert!(buf.is_ok(), "Should be able to allocate from SHM pool");
}
#[test]
#[serial]
fn test_shm_config_from_env_none() {
unsafe {
std::env::remove_var("ZENOH_SHM_ALLOC_SIZE");
std::env::remove_var("ZENOH_SHM_MESSAGE_SIZE_THRESHOLD");
}
let config = ShmConfig::from_env().expect("Should not error");
assert!(config.is_none(), "Should return None when no env vars set");
}
#[test]
#[serial]
fn test_shm_config_from_env_with_size() {
unsafe {
std::env::set_var("ZENOH_SHM_ALLOC_SIZE", "1048576"); std::env::remove_var("ZENOH_SHM_MESSAGE_SIZE_THRESHOLD");
}
let config = ShmConfig::from_env()
.expect("Should not error")
.expect("Should return Some when env var set");
assert_eq!(config.threshold(), DEFAULT_SHM_THRESHOLD);
unsafe {
std::env::remove_var("ZENOH_SHM_ALLOC_SIZE");
}
}
#[test]
#[serial]
fn test_shm_config_from_env_full() {
unsafe {
std::env::set_var("ZENOH_SHM_ALLOC_SIZE", "1048576"); std::env::set_var("ZENOH_SHM_MESSAGE_SIZE_THRESHOLD", "2048");
}
let config = ShmConfig::from_env()
.expect("Should not error")
.expect("Should return Some when env vars set");
assert_eq!(config.threshold(), 2048);
unsafe {
std::env::remove_var("ZENOH_SHM_ALLOC_SIZE");
std::env::remove_var("ZENOH_SHM_MESSAGE_SIZE_THRESHOLD");
}
}
}