use std::ops::{Deref, DerefMut};
use std::sync::{Arc, Mutex};
use crate::error::{Result, SaferRingError};
#[derive(Debug)]
pub enum BufferOwnership {
User(Box<[u8]>),
Kernel(Box<[u8]>, u64),
Returning,
}
#[derive(Debug)]
pub struct OwnedBuffer {
inner: Arc<Mutex<BufferOwnership>>,
size: usize,
generation: u64, }
impl OwnedBuffer {
pub fn new(size: usize) -> Self {
let buffer = vec![0u8; size].into_boxed_slice();
Self {
inner: Arc::new(Mutex::new(BufferOwnership::User(buffer))),
size,
generation: 0, }
}
pub fn from_slice(data: &[u8]) -> Self {
let buffer = data.to_vec().into_boxed_slice();
let size = buffer.len();
Self {
inner: Arc::new(Mutex::new(BufferOwnership::User(buffer))),
size,
generation: 0,
}
}
pub fn size(&self) -> usize {
self.size
}
pub fn generation(&self) -> u64 {
self.generation
}
pub fn try_access(&self) -> Option<BufferAccessGuard> {
let mut ownership = self.inner.lock().unwrap();
match &mut *ownership {
BufferOwnership::User(ref mut buf) => {
let buffer = std::mem::replace(buf, Box::new([]));
*ownership = BufferOwnership::Returning; Some(BufferAccessGuard {
buffer,
ownership: self.inner.clone(),
})
}
BufferOwnership::Kernel(_, _) => None, BufferOwnership::Returning => None, }
}
pub fn give_to_kernel(&self, submission_id: u64) -> Result<(*mut u8, usize)> {
let mut ownership = self.inner.lock().unwrap();
match &mut *ownership {
BufferOwnership::User(buf) => {
let ptr = buf.as_mut_ptr();
let len = buf.len();
let buffer = std::mem::replace(buf, Box::new([]));
*ownership = BufferOwnership::Kernel(buffer, submission_id);
Ok((ptr, len))
}
BufferOwnership::Kernel(_, existing_id) => {
Err(SaferRingError::Io(std::io::Error::new(
std::io::ErrorKind::WouldBlock,
format!("Buffer already owned by kernel (operation {existing_id})"),
)))
}
BufferOwnership::Returning => Err(SaferRingError::Io(std::io::Error::new(
std::io::ErrorKind::WouldBlock,
"Buffer ownership is currently transitioning",
))),
}
}
pub fn return_from_kernel(&self, submission_id: u64) {
let mut ownership = self.inner.lock().unwrap();
match &mut *ownership {
BufferOwnership::Kernel(buffer, id) if *id == submission_id => {
let buffer = std::mem::replace(buffer, Box::new([]));
*ownership = BufferOwnership::User(buffer);
}
BufferOwnership::Kernel(_, id) => {
panic!(
"Buffer owned by operation {id} but tried to return from operation {submission_id}"
);
}
_ => {
panic!("Tried to return buffer that isn't owned by kernel");
}
}
}
pub fn is_user_owned(&self) -> bool {
let ownership = self.inner.lock().unwrap();
matches!(*ownership, BufferOwnership::User(_))
}
pub fn kernel_owner(&self) -> Option<u64> {
let ownership = self.inner.lock().unwrap();
match *ownership {
BufferOwnership::Kernel(_, id) => Some(id),
_ => None,
}
}
pub fn clone_handle(&self) -> Self {
Self {
inner: self.inner.clone(),
size: self.size,
generation: self.generation,
}
}
pub fn as_ptr_and_len(&self) -> (*mut u8, usize) {
let ownership = self.inner.lock().unwrap();
match &*ownership {
BufferOwnership::User(buf) => (buf.as_ptr() as *mut u8, buf.len()),
BufferOwnership::Kernel(_, _) => (std::ptr::null_mut(), 0),
BufferOwnership::Returning => (std::ptr::null_mut(), 0),
}
}
}
pub struct BufferAccessGuard {
buffer: Box<[u8]>,
ownership: Arc<Mutex<BufferOwnership>>,
}
impl BufferAccessGuard {
pub fn len(&self) -> usize {
self.buffer.len()
}
pub fn is_empty(&self) -> bool {
self.buffer.is_empty()
}
pub fn as_ptr(&self) -> *const u8 {
self.buffer.as_ptr()
}
pub fn as_mut_ptr(&mut self) -> *mut u8 {
self.buffer.as_mut_ptr()
}
}
impl Drop for BufferAccessGuard {
fn drop(&mut self) {
let buffer = std::mem::replace(&mut self.buffer, Box::new([]));
let mut ownership = self.ownership.lock().unwrap();
*ownership = BufferOwnership::User(buffer);
}
}
impl Deref for BufferAccessGuard {
type Target = [u8];
fn deref(&self) -> &[u8] {
&self.buffer
}
}
impl DerefMut for BufferAccessGuard {
fn deref_mut(&mut self) -> &mut [u8] {
&mut self.buffer
}
}
pub trait SafeBuffer {
fn size(&self) -> usize;
fn give_to_kernel(&self, submission_id: u64) -> Result<(*mut u8, usize)>;
fn return_from_kernel(&self, submission_id: u64);
}
impl SafeBuffer for OwnedBuffer {
fn size(&self) -> usize {
self.size()
}
fn give_to_kernel(&self, submission_id: u64) -> Result<(*mut u8, usize)> {
self.give_to_kernel(submission_id)
}
fn return_from_kernel(&self, submission_id: u64) {
self.return_from_kernel(submission_id)
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_owned_buffer_creation() {
let buffer = OwnedBuffer::new(1024);
assert_eq!(buffer.size(), 1024);
assert!(buffer.is_user_owned());
assert_eq!(buffer.kernel_owner(), None);
}
#[test]
fn test_buffer_access_guard() {
let buffer = OwnedBuffer::new(1024);
{
let mut guard = buffer.try_access().unwrap();
assert_eq!(guard.len(), 1024);
guard[0] = 42;
guard[1] = 24;
}
let guard = buffer.try_access().unwrap();
assert_eq!(guard[0], 42);
assert_eq!(guard[1], 24);
}
#[test]
fn test_kernel_ownership_transfer() {
let buffer = OwnedBuffer::new(1024);
let (ptr, size) = buffer.give_to_kernel(123).unwrap();
assert!(!ptr.is_null());
assert_eq!(size, 1024);
assert!(!buffer.is_user_owned());
assert_eq!(buffer.kernel_owner(), Some(123));
assert!(buffer.try_access().is_none());
assert!(buffer.give_to_kernel(456).is_err());
}
#[test]
fn test_return_from_kernel() {
let buffer = OwnedBuffer::new(1024);
let (_ptr, _size) = buffer.give_to_kernel(123).unwrap();
buffer.return_from_kernel(123);
assert!(buffer.is_user_owned());
assert!(buffer.try_access().is_some());
}
#[test]
#[should_panic(
expected = "Buffer owned by operation 123 but tried to return from operation 456"
)]
fn test_mismatched_submission_id_panic() {
let buffer = OwnedBuffer::new(1024);
let (_ptr, _size) = buffer.give_to_kernel(123).unwrap();
buffer.return_from_kernel(456); }
#[test]
fn test_buffer_from_slice() {
let data = b"Hello, world!";
let buffer = OwnedBuffer::from_slice(data);
assert_eq!(buffer.size(), data.len());
let guard = buffer.try_access().unwrap();
assert_eq!(&*guard, data);
}
#[test]
fn test_clone_handle() {
let buffer = OwnedBuffer::new(1024);
let cloned = buffer.clone_handle();
assert_eq!(buffer.size(), cloned.size());
assert_eq!(buffer.generation(), cloned.generation());
buffer.give_to_kernel(123).unwrap();
assert!(!cloned.is_user_owned());
assert_eq!(cloned.kernel_owner(), Some(123));
}
}