use std::collections::HashSet;
use std::marker::PhantomData;
use std::os::unix::io::RawFd;
use std::pin::Pin;
use crate::error::{Result, SaferRingError};
use crate::ownership::OwnedBuffer;
#[cfg(target_os = "linux")]
#[cfg(test)]
mod tests;
pub struct Registry<'ring> {
registered_fds: Vec<Option<(RawFd, RegisteredFdInner)>>,
registered_buffers: Vec<Option<Pin<Box<[u8]>>>>,
fixed_files: Vec<Option<RawFd>>,
registered_buffer_slots: Vec<Option<OwnedBuffer>>,
fds_in_use: HashSet<u32>,
buffers_in_use: HashSet<u32>,
fixed_files_in_use: HashSet<u32>,
buffer_slots_in_use: HashSet<u32>,
#[allow(dead_code)]
is_registered: bool,
_phantom: PhantomData<&'ring ()>,
}
#[derive(Debug, Clone)]
struct RegisteredFdInner {
#[allow(dead_code)]
fd: RawFd,
#[allow(dead_code)]
in_use: bool,
}
#[derive(Debug)]
pub struct RegisteredFd {
index: u32,
fd: RawFd,
}
#[derive(Debug)]
pub struct RegisteredBuffer {
index: u32,
size: usize,
}
impl<'ring> Default for Registry<'ring> {
fn default() -> Self {
Self::new()
}
}
impl<'ring> Registry<'ring> {
pub fn new() -> Self {
Self {
registered_fds: Vec::new(),
registered_buffers: Vec::new(),
fixed_files: Vec::new(),
registered_buffer_slots: Vec::new(),
fds_in_use: HashSet::new(),
buffers_in_use: HashSet::new(),
fixed_files_in_use: HashSet::new(),
buffer_slots_in_use: HashSet::new(),
is_registered: false,
_phantom: PhantomData,
}
}
pub fn register_fd(&mut self, fd: RawFd) -> Result<RegisteredFd> {
if fd < 0 {
return Err(SaferRingError::Io(std::io::Error::new(
std::io::ErrorKind::InvalidInput,
"File descriptor must be non-negative",
)));
}
let index =
if let Some(empty_index) = self.registered_fds.iter().position(|slot| slot.is_none()) {
empty_index as u32
} else {
let index = self.registered_fds.len() as u32;
self.registered_fds.push(None);
index
};
let inner = RegisteredFdInner { fd, in_use: false };
self.registered_fds[index as usize] = Some((fd, inner));
Ok(RegisteredFd { index, fd })
}
pub fn register_buffer(&mut self, buffer: Pin<Box<[u8]>>) -> Result<RegisteredBuffer> {
if buffer.is_empty() {
return Err(SaferRingError::Io(std::io::Error::new(
std::io::ErrorKind::InvalidInput,
"Buffer cannot be empty",
)));
}
let size = buffer.len();
let index = if let Some(empty_index) = self
.registered_buffers
.iter()
.position(|slot| slot.is_none())
{
empty_index as u32
} else {
let index = self.registered_buffers.len() as u32;
self.registered_buffers.push(None);
index
};
self.registered_buffers[index as usize] = Some(buffer);
Ok(RegisteredBuffer { index, size })
}
pub fn unregister_fd(&mut self, registered_fd: RegisteredFd) -> Result<()> {
let index = registered_fd.index as usize;
if index >= self.registered_fds.len() {
return Err(SaferRingError::NotRegistered);
}
let slot = &mut self.registered_fds[index];
let Some((stored_fd, _)) = slot.as_ref() else {
return Err(SaferRingError::NotRegistered);
};
if *stored_fd != registered_fd.fd {
return Err(SaferRingError::NotRegistered);
}
if self.fds_in_use.contains(®istered_fd.index) {
return Err(SaferRingError::BufferInFlight);
}
*slot = None;
Ok(())
}
pub fn unregister_buffer(
&mut self,
registered_buffer: RegisteredBuffer,
) -> Result<Pin<Box<[u8]>>> {
let index = registered_buffer.index as usize;
if index >= self.registered_buffers.len() {
return Err(SaferRingError::NotRegistered);
}
let slot = &mut self.registered_buffers[index];
let Some(buffer) = slot.as_ref() else {
return Err(SaferRingError::NotRegistered);
};
if buffer.len() != registered_buffer.size {
return Err(SaferRingError::NotRegistered);
}
if self.buffers_in_use.contains(®istered_buffer.index) {
return Err(SaferRingError::BufferInFlight);
}
let buffer = slot.take().unwrap();
Ok(buffer)
}
#[allow(dead_code)]
pub(crate) fn mark_fd_in_use(&mut self, registered_fd: &RegisteredFd) -> Result<()> {
let index = registered_fd.index as usize;
if index >= self.registered_fds.len() || self.registered_fds[index].is_none() {
return Err(SaferRingError::NotRegistered);
}
self.fds_in_use.insert(registered_fd.index);
Ok(())
}
#[allow(dead_code)]
pub(crate) fn mark_fd_not_in_use(&mut self, registered_fd: &RegisteredFd) {
self.fds_in_use.remove(®istered_fd.index);
}
#[allow(dead_code)]
pub(crate) fn mark_buffer_in_use(
&mut self,
registered_buffer: &RegisteredBuffer,
) -> Result<()> {
let index = registered_buffer.index as usize;
if index >= self.registered_buffers.len() || self.registered_buffers[index].is_none() {
return Err(SaferRingError::NotRegistered);
}
self.buffers_in_use.insert(registered_buffer.index);
Ok(())
}
#[allow(dead_code)]
pub(crate) fn mark_buffer_not_in_use(&mut self, registered_buffer: &RegisteredBuffer) {
self.buffers_in_use.remove(®istered_buffer.index);
}
pub fn fd_count(&self) -> usize {
self.registered_fds
.iter()
.filter(|slot| slot.is_some())
.count()
}
pub fn buffer_count(&self) -> usize {
self.registered_buffers
.iter()
.filter(|slot| slot.is_some())
.count()
}
pub fn fds_in_use_count(&self) -> usize {
self.fds_in_use.len()
}
pub fn buffers_in_use_count(&self) -> usize {
self.buffers_in_use.len()
}
pub fn is_fd_registered(&self, registered_fd: &RegisteredFd) -> bool {
let index = registered_fd.index as usize;
index < self.registered_fds.len()
&& self.registered_fds[index]
.as_ref()
.map(|(fd, _)| *fd == registered_fd.fd)
.unwrap_or(false)
}
pub fn is_buffer_registered(&self, registered_buffer: &RegisteredBuffer) -> bool {
let index = registered_buffer.index as usize;
index < self.registered_buffers.len()
&& self.registered_buffers[index]
.as_ref()
.map(|buffer| buffer.len() == registered_buffer.size)
.unwrap_or(false)
}
#[allow(dead_code)]
pub(crate) fn get_raw_fd(&self, registered_fd: &RegisteredFd) -> Result<RawFd> {
let index = registered_fd.index as usize;
if index >= self.registered_fds.len() {
return Err(SaferRingError::NotRegistered);
}
match &self.registered_fds[index] {
Some((fd, _)) if *fd == registered_fd.fd => Ok(*fd),
_ => Err(SaferRingError::NotRegistered),
}
}
#[allow(dead_code)]
pub(crate) fn get_buffer(
&self,
registered_buffer: &RegisteredBuffer,
) -> Result<&Pin<Box<[u8]>>> {
let index = registered_buffer.index as usize;
if index >= self.registered_buffers.len() {
return Err(SaferRingError::NotRegistered);
}
match &self.registered_buffers[index] {
Some(buffer) if buffer.len() == registered_buffer.size => Ok(buffer),
_ => Err(SaferRingError::NotRegistered),
}
}
pub fn register_fixed_files(&mut self, fds: Vec<RawFd>) -> Result<Vec<FixedFile>> {
if fds.is_empty() {
return Ok(Vec::new());
}
for &fd in &fds {
if fd < 0 {
return Err(SaferRingError::Io(std::io::Error::new(
std::io::ErrorKind::InvalidInput,
format!("Invalid file descriptor: {fd}"),
)));
}
}
self.fixed_files.clear();
self.fixed_files_in_use.clear();
let mut fixed_files = Vec::new();
for (index, &fd) in fds.iter().enumerate() {
self.fixed_files.push(Some(fd));
fixed_files.push(FixedFile {
index: index as u32,
fd,
});
}
Ok(fixed_files)
}
pub fn unregister_fixed_files(&mut self) -> Result<()> {
if !self.fixed_files_in_use.is_empty() {
return Err(SaferRingError::BufferInFlight);
}
self.fixed_files.clear();
Ok(())
}
pub fn fixed_file_count(&self) -> usize {
self.fixed_files
.iter()
.filter(|slot| slot.is_some())
.count()
}
pub fn fixed_files_in_use_count(&self) -> usize {
self.fixed_files_in_use.len()
}
pub fn register_buffer_slots(
&mut self,
buffers: Vec<OwnedBuffer>,
) -> Result<Vec<RegisteredBufferSlot>> {
if buffers.is_empty() {
return Ok(Vec::new());
}
self.registered_buffer_slots.clear();
self.buffer_slots_in_use.clear();
let mut buffer_slots = Vec::new();
for (index, buffer) in buffers.into_iter().enumerate() {
let size = buffer.size();
self.registered_buffer_slots.push(Some(buffer));
buffer_slots.push(RegisteredBufferSlot {
index: index as u32,
size,
in_use: false,
});
}
Ok(buffer_slots)
}
pub fn unregister_buffer_slots(&mut self) -> Result<Vec<OwnedBuffer>> {
if !self.buffer_slots_in_use.is_empty() {
return Err(SaferRingError::BufferInFlight);
}
let buffers = self.registered_buffer_slots.drain(..).flatten().collect();
Ok(buffers)
}
pub fn buffer_slot_count(&self) -> usize {
self.registered_buffer_slots
.iter()
.filter(|slot| slot.is_some())
.count()
}
pub fn buffer_slots_in_use_count(&self) -> usize {
self.buffer_slots_in_use.len()
}
#[allow(dead_code)]
pub(crate) fn mark_fixed_file_in_use(&mut self, fixed_file: &FixedFile) -> Result<()> {
let index = fixed_file.index as usize;
if index >= self.fixed_files.len() || self.fixed_files[index].is_none() {
return Err(SaferRingError::NotRegistered);
}
self.fixed_files_in_use.insert(fixed_file.index);
Ok(())
}
#[allow(dead_code)]
pub(crate) fn mark_fixed_file_not_in_use(&mut self, fixed_file: &FixedFile) {
self.fixed_files_in_use.remove(&fixed_file.index);
}
#[allow(dead_code)]
pub(crate) fn mark_buffer_slot_in_use(
&mut self,
buffer_slot: &RegisteredBufferSlot,
) -> Result<()> {
let index = buffer_slot.index as usize;
if index >= self.registered_buffer_slots.len()
|| self.registered_buffer_slots[index].is_none()
{
return Err(SaferRingError::NotRegistered);
}
self.buffer_slots_in_use.insert(buffer_slot.index);
Ok(())
}
#[allow(dead_code)]
pub(crate) fn mark_buffer_slot_not_in_use(&mut self, buffer_slot: &RegisteredBufferSlot) {
self.buffer_slots_in_use.remove(&buffer_slot.index);
}
#[allow(dead_code)]
pub(crate) fn get_fixed_file_fd(&self, fixed_file: &FixedFile) -> Result<RawFd> {
let index = fixed_file.index as usize;
if index >= self.fixed_files.len() {
return Err(SaferRingError::NotRegistered);
}
match &self.fixed_files[index] {
Some(fd) if *fd == fixed_file.fd => Ok(*fd),
_ => Err(SaferRingError::NotRegistered),
}
}
#[allow(dead_code)]
pub(crate) fn get_buffer_slot(
&self,
buffer_slot: &RegisteredBufferSlot,
) -> Result<&OwnedBuffer> {
let index = buffer_slot.index as usize;
if index >= self.registered_buffer_slots.len() {
return Err(SaferRingError::NotRegistered);
}
match &self.registered_buffer_slots[index] {
Some(buffer) if buffer.size() == buffer_slot.size => Ok(buffer),
_ => Err(SaferRingError::NotRegistered),
}
}
}
impl<'ring> Drop for Registry<'ring> {
fn drop(&mut self) {
let total_in_use = self.fds_in_use.len()
+ self.buffers_in_use.len()
+ self.fixed_files_in_use.len()
+ self.buffer_slots_in_use.len();
if total_in_use > 0 {
panic!(
"Registry dropped with resources in use: {} fds, {} buffers, {} fixed_files, {} buffer_slots",
self.fds_in_use.len(),
self.buffers_in_use.len(),
self.fixed_files_in_use.len(),
self.buffer_slots_in_use.len()
);
}
}
}
impl RegisteredFd {
pub fn index(&self) -> u32 {
self.index
}
pub fn raw_fd(&self) -> RawFd {
self.fd
}
}
impl RegisteredBuffer {
pub fn index(&self) -> u32 {
self.index
}
pub fn size(&self) -> usize {
self.size
}
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct FixedFile {
index: u32,
fd: RawFd,
}
#[derive(Debug)]
pub struct RegisteredBufferSlot {
index: u32,
size: usize,
in_use: bool,
}
impl FixedFile {
pub fn index(&self) -> u32 {
self.index
}
pub fn raw_fd(&self) -> RawFd {
self.fd
}
}
impl RegisteredBufferSlot {
pub fn index(&self) -> u32 {
self.index
}
pub fn size(&self) -> usize {
self.size
}
pub fn is_in_use(&self) -> bool {
self.in_use
}
}