use alloc::boxed::Box;
use core::future::Future;
use core::marker::PhantomData;
use core::pin::Pin;
struct Guard<A, R: FnOnce(A)> {
resource: Option<A>,
release: Option<R>,
}
impl<A, R: FnOnce(A)> Drop for Guard<A, R> {
fn drop(&mut self) {
if let (Some(resource), Some(release)) = (self.resource.take(), self.release.take()) {
release(resource);
}
}
}
pub struct Res<R, A> {
acquire: Box<dyn FnOnce() -> A + Send>,
release: Box<dyn FnOnce(A) + Send>,
_resource: PhantomData<R>,
}
impl<R, A> Res<R, A> {
#[inline]
pub fn make<Acq, Rel>(acquire: Acq, release: Rel) -> Self
where
Acq: FnOnce() -> A + Send + 'static,
Rel: FnOnce(A) + Send + 'static,
{
Res {
acquire: Box::new(acquire),
release: Box::new(release),
_resource: PhantomData,
}
}
#[inline]
pub fn use_res<B, F>(self, f: F) -> B
where
F: FnOnce(&A) -> B,
{
let resource = (self.acquire)();
let guard = Guard {
resource: Some(resource),
release: Some(self.release),
};
let result = f(guard.resource.as_ref().unwrap());
drop(guard);
result
}
#[inline]
pub fn use_res_mut<B, F>(self, f: F) -> B
where
F: FnOnce(&mut A) -> B,
{
let resource = (self.acquire)();
let mut guard = Guard {
resource: Some(resource),
release: Some(self.release),
};
let result = f(guard.resource.as_mut().unwrap());
drop(guard);
result
}
#[inline]
pub fn fmap<B, F, G>(self, transform: F, inverse: G) -> Res<R, B>
where
F: FnOnce(A) -> B + Send + 'static,
G: FnOnce(B) -> A + Send + 'static,
A: 'static,
{
let acquire = self.acquire;
let release = self.release;
Res {
acquire: Box::new(move || transform((acquire)())),
release: Box::new(move |b| (release)(inverse(b))),
_resource: PhantomData,
}
}
}
impl<R, A: 'static + Send> Res<R, A> {
#[inline]
pub fn zip<B: 'static + Send>(self, other: Res<R, B>) -> Res<R, (A, B)> {
let acquire1 = self.acquire;
let release1 = self.release;
let acquire2 = other.acquire;
let release2 = other.release;
Res {
acquire: Box::new(move || ((acquire1)(), (acquire2)())),
release: Box::new(move |(a, b)| {
(release1)(a);
(release2)(b);
}),
_resource: PhantomData,
}
}
}
type AsyncReleaseFn<A> = Box<dyn FnOnce(A) -> Pin<Box<dyn Future<Output = ()> + Send>> + Send>;
pub struct ResAsync<R, A> {
acquire: Pin<Box<dyn Future<Output = A> + Send>>,
release: AsyncReleaseFn<A>,
_resource: PhantomData<R>,
}
impl<R, A> ResAsync<R, A> {
#[inline]
pub fn make<Acq, Rel, RelFut>(acquire: Acq, release: Rel) -> Self
where
Acq: Future<Output = A> + Send + 'static,
Rel: FnOnce(A) -> RelFut + Send + 'static,
RelFut: Future<Output = ()> + Send + 'static,
{
ResAsync {
acquire: Box::pin(acquire),
release: Box::new(move |a| Box::pin(release(a))),
_resource: PhantomData,
}
}
pub async fn use_res_async<B, F, Fut>(self, f: F) -> B
where
F: FnOnce(A) -> Fut,
Fut: Future<Output = (A, B)>,
{
let resource = self.acquire.await;
let (resource, result) = f(resource).await;
(self.release)(resource).await;
result
}
}
#[inline]
pub fn bracket<A, B, Acquire, Use, Release>(acquire: Acquire, use_fn: Use, release: Release) -> B
where
Acquire: FnOnce() -> A,
Use: FnOnce(&A) -> B,
Release: FnOnce(A),
{
let resource = acquire();
let guard = Guard {
resource: Some(resource),
release: Some(release),
};
let result = use_fn(guard.resource.as_ref().unwrap());
drop(guard);
result
}
#[cfg(feature = "std")]
#[inline]
pub fn bracket_safe<A, B, Acquire, Use, Release>(
acquire: Acquire,
use_fn: Use,
release: Release,
) -> Result<B, Box<dyn std::any::Any + Send>>
where
A: std::panic::UnwindSafe,
Acquire: FnOnce() -> A,
Use: FnOnce(&A) -> B + std::panic::UnwindSafe,
Release: FnOnce(A),
{
let resource = acquire();
let result = std::panic::catch_unwind(std::panic::AssertUnwindSafe(|| use_fn(&resource)));
release(resource);
result
}
#[cfg(test)]
mod tests {
use super::*;
use alloc::format;
use alloc::string::{String, ToString};
use alloc::sync::Arc;
use alloc::vec::Vec;
use core::sync::atomic::{AtomicBool, Ordering};
#[test]
fn test_res_basic() {
let released = Arc::new(AtomicBool::new(false));
let released_clone = Arc::clone(&released);
let res = Res::<(), i32>::make(
|| 42,
move |_| {
released_clone.store(true, Ordering::SeqCst);
},
);
let result = res.use_res(|&n| n * 2);
assert_eq!(result, 84);
assert!(released.load(Ordering::SeqCst));
}
#[test]
fn test_res_use_mut() {
let res = Res::<(), Vec<i32>>::make(Vec::new, |_| {});
let result = res.use_res_mut(|v| {
v.push(1);
v.push(2);
v.len()
});
assert_eq!(result, 2);
}
#[test]
fn test_res_zip() {
let res1 = Res::<(), i32>::make(|| 1, |_| {});
let res2 = Res::<(), String>::make(|| "hello".to_string(), |_| {});
let combined = res1.zip(res2);
let result = combined.use_res(|(n, s)| format!("{n}: {s}"));
assert_eq!(result, "1: hello");
}
#[test]
fn test_bracket() {
let released = Arc::new(AtomicBool::new(false));
let released_clone = Arc::clone(&released);
let result = bracket(
|| 42,
|&n| n * 2,
move |_| {
released_clone.store(true, Ordering::SeqCst);
},
);
assert_eq!(result, 84);
assert!(released.load(Ordering::SeqCst));
}
#[test]
fn test_res_fmap() {
let res = Res::<(), i32>::make(|| 42, |_| {});
let mapped = res.fmap(
|n| n.to_string(),
|s| {
s.parse()
.expect("fmap inverse should parse '42' back to i32")
},
);
let result = mapped.use_res(std::string::String::len);
assert_eq!(result, 2); }
#[test]
#[cfg(feature = "std")]
fn test_res_panic_safety() {
use std::panic::{AssertUnwindSafe, catch_unwind};
let released = Arc::new(AtomicBool::new(false));
let released_clone = Arc::clone(&released);
let res = Res::<(), i32>::make(
|| 42,
move |_| {
released_clone.store(true, Ordering::SeqCst);
},
);
let _ = catch_unwind(AssertUnwindSafe(|| {
res.use_res(|_| {
panic!("Oops");
});
}));
assert!(
released.load(Ordering::SeqCst),
"Resource should be released on panic"
);
}
}