use std::{
any::{Any, TypeId},
cell::{Cell, RefCell},
collections::{BTreeMap, BTreeSet},
fmt,
rc::Rc,
task::Poll,
};
#[cfg(test)]
mod tests;
use futures::{
future::{LocalBoxFuture, poll_fn},
task::AtomicWaker,
};
use lenso_kernel::RuntimeFailure;
use serde_json::Value;
pub trait NativeGenerationResource: fmt::Debug + 'static {
fn revoke(&self) -> LocalBoxFuture<'static, Result<(), RuntimeFailure>>;
}
struct Entry {
type_id: TypeId,
configuration: Value,
value: Rc<dyn Any>,
revoke: Rc<dyn Fn() -> LocalBoxFuture<'static, Result<(), RuntimeFailure>>>,
}
#[derive(Default)]
struct State {
closed: Cell<bool>,
leases: Cell<usize>,
wake: AtomicWaker,
entries: RefCell<BTreeMap<(String, String), Entry>>,
order: RefCell<Vec<(String, String)>>,
constructing: RefCell<BTreeSet<(String, String)>>,
retirement: RefCell<Option<LocalBoxFuture<'static, Result<(), RuntimeFailure>>>>,
outcome: RefCell<Option<Result<(), RuntimeFailure>>>,
retiring: Cell<bool>,
}
struct PollGuard(Rc<State>);
impl Drop for PollGuard {
fn drop(&mut self) {
self.0.retiring.set(false);
}
}
struct Reservation {
state: Rc<State>,
key: (String, String),
}
impl Drop for Reservation {
fn drop(&mut self) {
self.state.constructing.borrow_mut().remove(&self.key);
self.state.wake.wake();
}
}
pub struct NativeGenerationScope {
state: Rc<State>,
}
impl fmt::Debug for NativeGenerationScope {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("NativeGenerationScope")
.field("closed", &self.state.closed.get())
.field("leases", &self.state.leases.get())
.finish_non_exhaustive()
}
}
impl Default for NativeGenerationScope {
fn default() -> Self {
Self::new()
}
}
impl NativeGenerationScope {
pub fn new() -> Self {
Self {
state: Rc::default(),
}
}
pub fn context(&self) -> NativeGenerationContext {
NativeGenerationContext {
state: self.state.clone(),
}
}
pub fn close(&self) {
self.state.closed.set(true);
self.state.wake.wake();
}
pub async fn retire(&self) -> Result<(), RuntimeFailure> {
self.close();
if self.state.retiring.replace(true) {
return Err(invalid("generation retirement already being polled"));
}
let _guard = PollGuard(self.state.clone());
poll_fn(|cx| {
if let Some(result) = self.state.outcome.borrow().as_ref() {
return Poll::Ready(result.clone());
}
self.state.wake.register(cx.waker());
if self.state.leases.get() != 0 || !self.state.constructing.borrow().is_empty() {
return Poll::Pending;
}
if self.state.retirement.borrow().is_none() {
let resources: Vec<_> = self
.state
.order
.borrow()
.iter()
.map(|key| self.state.entries.borrow()[key].revoke.clone())
.collect();
*self.state.retirement.borrow_mut() = Some(Box::pin(async move {
let mut failure = None;
for resource in resources.into_iter().rev() {
if let Err(error) = resource().await
&& failure.is_none()
{
failure = Some(error);
}
}
failure.map_or(Ok(()), Err)
}));
}
let result = self
.state
.retirement
.borrow_mut()
.as_mut()
.expect("retirement initialized")
.as_mut()
.poll(cx);
if let Poll::Ready(result) = result {
if result.is_ok() {
self.state.entries.borrow_mut().clear();
self.state.order.borrow_mut().clear();
}
*self.state.outcome.borrow_mut() = Some(result.clone());
Poll::Ready(result)
} else {
Poll::Pending
}
})
.await
}
}
impl Drop for NativeGenerationScope {
fn drop(&mut self) {
self.close();
if !self.state.entries.borrow().is_empty() || !self.state.constructing.borrow().is_empty() {
std::mem::forget(self.state.clone());
}
}
}
#[derive(Clone)]
pub struct NativeGenerationContext {
state: Rc<State>,
}
impl fmt::Debug for NativeGenerationContext {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("NativeGenerationContext")
.finish_non_exhaustive()
}
}
impl NativeGenerationContext {
pub fn resolve<T: NativeGenerationResource>(
&self,
owner: &str,
reference: &str,
configuration: &Value,
create: impl FnOnce() -> Result<T, RuntimeFailure>,
) -> Result<NativeGenerationHandle<T>, RuntimeFailure> {
if self.state.closed.get() || !valid_reference(owner) || !valid_reference(reference) {
return Err(invalid("invalid or closed generation authority reference"));
}
let key = (owner.to_owned(), reference.to_owned());
if let Some(entry) = self.state.entries.borrow().get(&key) {
if entry.type_id != TypeId::of::<T>() || &entry.configuration != configuration {
return Err(invalid(
"generation authority owner/reference has conflicting type or configuration",
));
}
return Ok(NativeGenerationHandle {
value: entry
.value
.clone()
.downcast::<T>()
.map_err(|_| invalid("generation authority type mismatch"))?,
state: self.state.clone(),
});
}
if !self.state.constructing.borrow_mut().insert(key.clone()) {
return Err(invalid(
"generation authority constructor reentered its own reference",
));
}
let _reservation = Reservation {
state: self.state.clone(),
key: key.clone(),
};
let value = Rc::new(create()?);
let revoke = value.clone();
self.state.order.borrow_mut().push(key.clone());
self.state.entries.borrow_mut().insert(
key,
Entry {
type_id: TypeId::of::<T>(),
configuration: configuration.clone(),
value: value.clone(),
revoke: Rc::new(move || revoke.revoke()),
},
);
if self.state.closed.get() {
return Err(invalid(
"generation authority closed during construction; rollback required",
));
}
Ok(NativeGenerationHandle {
value,
state: self.state.clone(),
})
}
}
pub struct NativeGenerationHandle<T> {
value: Rc<T>,
state: Rc<State>,
}
impl<T> Clone for NativeGenerationHandle<T> {
fn clone(&self) -> Self {
Self {
value: self.value.clone(),
state: self.state.clone(),
}
}
}
impl<T> fmt::Debug for NativeGenerationHandle<T> {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("NativeGenerationHandle")
.finish_non_exhaustive()
}
}
impl<T> NativeGenerationHandle<T> {
pub fn acquire(&self) -> Result<NativeGenerationLease<T>, RuntimeFailure> {
if self.state.closed.get() {
return Err(invalid("generation authority admission is closed"));
}
self.state.leases.set(
self.state
.leases
.get()
.checked_add(1)
.ok_or_else(|| invalid("generation authority lease limit"))?,
);
Ok(NativeGenerationLease {
value: self.value.clone(),
state: self.state.clone(),
})
}
}
pub struct NativeGenerationLease<T> {
value: Rc<T>,
state: Rc<State>,
}
impl<T> fmt::Debug for NativeGenerationLease<T> {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("NativeGenerationLease")
.finish_non_exhaustive()
}
}
impl<T> std::ops::Deref for NativeGenerationLease<T> {
type Target = T;
fn deref(&self) -> &T {
&self.value
}
}
impl<T> Drop for NativeGenerationLease<T> {
fn drop(&mut self) {
self.state.leases.set(self.state.leases.get() - 1);
self.state.wake.wake();
}
}
fn invalid(detail: &str) -> RuntimeFailure {
RuntimeFailure::InvalidResolvedPlan {
detail: detail.into(),
}
}
fn valid_reference(reference: &str) -> bool {
!reference.is_empty()
&& reference.len() <= 128
&& reference
.bytes()
.all(|b| b.is_ascii_alphanumeric() || matches!(b, b'.' | b'_' | b'-'))
}