use std::sync::Arc;
#[cfg(test)]
use std::sync::atomic::AtomicUsize;
use std::sync::atomic::{AtomicBool, Ordering};
use std::time::{Duration, Instant};
use bytes::Bytes;
use tokio::sync::Notify;
use tokio_util::sync::CancellationToken;
use velo_ext::ShutdownState;
use super::RegistryShared;
use super::arena::RemoteRef;
use super::backend::RdmaError;
pub(crate) struct RegionInner {
pub(super) id: u64,
pub(super) generation: u64,
pub(super) backend_region_id: u64,
pub(super) ptr: usize,
pub(super) len: usize,
pub(super) packed_key: Bytes,
pub(super) effective_addr: u64,
pub(super) effective_len: u64,
pub(super) in_flight: ShutdownState,
deregistered: AtomicBool,
copy_gate: parking_lot::RwLock<()>,
dereg_notify: Notify,
pub(super) dereg_lock: tokio::sync::Mutex<()>,
pub(super) owned: parking_lot::Mutex<Option<Box<[u8]>>>,
charged: u64,
pub(super) shutdown: CancellationToken,
}
pub(super) struct RegionParts {
pub id: u64,
pub generation: u64,
pub backend_region_id: u64,
pub ptr: usize,
pub len: usize,
pub packed_key: Bytes,
pub effective_addr: u64,
pub effective_len: u64,
pub owned: Option<Box<[u8]>>,
pub charged: u64,
pub shutdown: CancellationToken,
}
impl RegionInner {
pub(super) fn new(parts: RegionParts) -> Self {
Self {
id: parts.id,
generation: parts.generation,
backend_region_id: parts.backend_region_id,
ptr: parts.ptr,
len: parts.len,
packed_key: parts.packed_key,
effective_addr: parts.effective_addr,
effective_len: parts.effective_len,
in_flight: ShutdownState::new(),
deregistered: AtomicBool::new(false),
copy_gate: parking_lot::RwLock::new(()),
dereg_notify: Notify::new(),
dereg_lock: tokio::sync::Mutex::new(()),
owned: parking_lot::Mutex::new(parts.owned),
charged: parts.charged,
shutdown: parts.shutdown,
}
}
pub(super) fn is_deregistered(&self) -> bool {
self.deregistered.load(Ordering::SeqCst)
}
pub(super) fn latch_deregistered(&self) {
let _closing = self.copy_gate.write();
self.deregistered.store(true, Ordering::SeqCst);
self.dereg_notify.notify_waiters();
}
pub(super) fn with_live<R>(&self, read: impl FnOnce() -> R) -> Option<R> {
let _open = self.copy_gate.read();
if self.is_deregistered() {
return None;
}
Some(read())
}
pub(super) fn charged(&self) -> u64 {
self.charged
}
pub(super) async fn wait_deregistered(&self) {
loop {
let notified = self.dereg_notify.notified();
if self.is_deregistered() {
return;
}
notified.await;
}
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
#[must_use = "a DrainTimedOut deregistration released the memory without waiting for in-flight work"]
pub enum Deregistered {
Drained,
DrainTimedOut,
}
#[cfg(test)]
pub(crate) static LEAKED_BUFFERS: AtomicUsize = AtomicUsize::new(0);
impl Drop for RegionInner {
fn drop(&mut self) {
let Some(buffer) = self.owned.get_mut().take() else {
return;
};
if self.is_deregistered() {
return;
}
let bytes = buffer.len();
let region = self.id;
let _ = Box::leak(buffer);
#[cfg(test)]
LEAKED_BUFFERS.fetch_add(1, Ordering::SeqCst);
tracing::error!(
region,
bytes,
"rdma: an owned registration was dropped without a confirmed deregistration; \
leaking its buffer rather than freeing memory the backend may still have \
pinned. The registry was torn down without `RdmaRegistry::shutdown`."
);
}
}
pub(super) async fn deregister(
shared: &Arc<RegistryShared>,
inner: &Arc<RegionInner>,
budget: Duration,
) -> Result<Deregistered, RdmaError> {
let deadline = Instant::now() + budget;
let remaining = |deadline: Instant| deadline.saturating_duration_since(Instant::now());
let Ok(_lock) = tokio::time::timeout(remaining(deadline), inner.dereg_lock.lock()).await else {
return Err(RdmaError::Timeout);
};
if inner.is_deregistered() {
return Ok(Deregistered::Drained);
}
inner.in_flight.begin_drain();
let drained = tokio::time::timeout(remaining(deadline), inner.in_flight.wait_for_drain())
.await
.is_ok();
if !drained {
let waiting = inner.in_flight.in_flight_count();
let region = inner.id;
tracing::warn!(
region,
waiting,
"rdma: region drain timed out; unmapping anyway"
);
}
let unmap = shared.backend.unmap(inner.backend_region_id);
let outcome = match tokio::time::timeout(remaining(deadline), unmap).await {
Ok(result) => result,
Err(_) => Err(RdmaError::Timeout),
};
match outcome {
Ok(()) => {
shared.forget_region(inner);
inner.latch_deregistered();
if drained {
Ok(Deregistered::Drained)
} else {
Ok(Deregistered::DrainTimedOut)
}
}
Err(e) => Err(e),
}
}
#[must_use = "the registration lasts as long as this guard; dropping it immediately starts a background deregistration"]
pub struct RegionGuard {
inner: Arc<RegionInner>,
shared: Arc<RegistryShared>,
}
impl RegionGuard {
pub(super) fn new(inner: Arc<RegionInner>, shared: Arc<RegistryShared>) -> Self {
Self { inner, shared }
}
pub fn addr(&self) -> u64 {
self.inner.ptr as u64
}
pub fn len(&self) -> u64 {
self.inner.len as u64
}
pub fn is_empty(&self) -> bool {
self.inner.len == 0
}
pub fn generation(&self) -> u64 {
self.inner.generation
}
pub fn effective_range(&self) -> (u64, u64) {
(self.inner.effective_addr, self.inner.effective_len)
}
pub fn is_shutting_down(&self) -> bool {
self.inner.shutdown.is_cancelled()
}
pub async fn shutdown_initiated(&self) {
self.inner.shutdown.cancelled().await;
}
pub fn is_deregistered(&self) -> bool {
self.inner.is_deregistered()
}
pub async fn deregistered(&self) {
self.inner.wait_deregistered().await;
}
pub fn watch(&self) -> RegionWatch {
RegionWatch {
inner: Arc::clone(&self.inner),
}
}
pub(crate) fn remote(&self) -> RemoteRef {
RemoteRef {
addr: self.inner.ptr as u64,
len: self.inner.len as u64,
packed_key: self.inner.packed_key.clone(),
generation: self.inner.generation,
}
}
pub(crate) fn in_flight(&self) -> &ShutdownState {
&self.inner.in_flight
}
pub async fn unregister(self, timeout: Duration) -> Result<Deregistered, RdmaError> {
deregister(&self.shared, &self.inner, timeout).await
}
pub async fn unregister_owned(
self,
timeout: Duration,
) -> Result<(Box<[u8]>, Deregistered), RdmaError> {
if self.inner.owned.lock().is_none() {
return Err(RdmaError::NotOwned);
}
let outcome = deregister(&self.shared, &self.inner, timeout).await?;
let taken = self.inner.owned.lock().take();
taken
.map(|buffer| (buffer, outcome))
.ok_or(RdmaError::NotOwned)
}
}
impl Drop for RegionGuard {
fn drop(&mut self) {
if self.inner.is_deregistered() {
return;
}
let region = self.inner.id;
let bytes = self.inner.len;
tracing::warn!(
region,
bytes,
"rdma: RegionGuard dropped before deregistration; deregistering in the \
background. The memory stays pinned until that finishes, so it is not safe \
to free yet — await `deregistered()` on a watch, or velo shutdown."
);
let shared = Arc::clone(&self.shared);
let inner = Arc::clone(&self.inner);
let budget = shared.cfg.drop_dereg_timeout;
shared.runtime.clone().spawn(async move {
if let Err(e) = deregister(&shared, &inner, budget).await {
tracing::error!(
region,
error = %e,
"rdma: background deregistration did not confirm; the region stays \
registered until velo shuts down"
);
}
});
}
}
impl std::fmt::Debug for RegionGuard {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("RegionGuard")
.field("region", &self.inner.id)
.field("addr", &self.inner.ptr)
.field("len", &self.inner.len)
.field("generation", &self.inner.generation)
.field("deregistered", &self.inner.is_deregistered())
.finish()
}
}
#[derive(Clone)]
pub struct RegionWatch {
inner: Arc<RegionInner>,
}
impl std::fmt::Debug for RegionWatch {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("RegionWatch")
.field("region", &self.inner.id)
.field("deregistered", &self.inner.is_deregistered())
.finish()
}
}
impl RegionWatch {
pub fn is_shutting_down(&self) -> bool {
self.inner.shutdown.is_cancelled()
}
pub async fn shutdown_initiated(&self) {
self.inner.shutdown.cancelled().await;
}
pub fn is_deregistered(&self) -> bool {
self.inner.is_deregistered()
}
pub async fn deregistered(&self) {
self.inner.wait_deregistered().await;
}
pub(crate) fn with_live<R>(&self, read: impl FnOnce() -> R) -> Option<R> {
self.inner.with_live(read)
}
#[cfg(test)]
pub(crate) fn for_test(inner: Arc<RegionInner>) -> Self {
Self { inner }
}
#[cfg(test)]
pub(crate) fn latch_for_test(&self) {
self.inner.latch_deregistered();
}
}