mod app_context;
mod id;
mod request_context;
mod tracking;
use std::{any::Any, panic::Location, sync::Arc};
pub use app_context::*;
pub use id::*;
pub use request_context::*;
pub(crate) use tracking::*;
pub use crate::memoize::MemoizeAsRef;
use crate::{
abort::AbortStore,
identity::{AmbiguousIdentityError, Identity, IdentityKey, SiteKey},
memoize::MemoizeCache,
};
#[derive(Debug, Default, Clone)]
pub struct Cx {
state: Arc<CxState>,
identity: Identity,
}
impl Cx {
#[must_use]
pub fn new(app_context: Arc<AppContext>) -> Self {
Self::from_parts(app_context, RequestContext::new())
}
fn from_parts(app_context: Arc<AppContext>, request_context: RequestContext) -> Self {
Self {
state: Arc::new(CxState {
shared: Arc::new(RequestShared {
id: CxId::new(),
app_context,
memoize_cache: MemoizeCache::new(),
abort_store: AbortStore::new(),
}),
request_context: Arc::new(request_context),
tracker: None,
}),
identity: Identity::ROOT,
}
}
#[inline]
#[must_use]
pub fn id(&self) -> CxId {
self.state.shared.id
}
#[must_use]
#[track_caller]
pub fn keyed(&self, key: impl IdentityKey) -> Self {
let site = SiteKey::from_location(Location::caller());
Self {
identity: self.identity.keyed_child(site, key),
..self.clone()
}
}
#[inline]
pub(crate) fn request_context(&self) -> &RequestContext {
&self.state.request_context
}
#[inline]
pub(crate) fn tracker(&self) -> Option<&ContextTracker> {
self.state.tracker.as_deref()
}
#[must_use]
pub fn with<T>(&self, value: T) -> Cx
where
T: Any + Send + Sync,
{
let mut request_context = (*self.state.request_context).clone();
request_context.insert(value);
self.scope(request_context)
}
#[must_use]
pub fn with_many<V>(&self, values: V) -> Cx
where
V: ContextValues,
{
let mut request_context = (*self.state.request_context).clone();
values.install(&mut request_context);
self.scope(request_context)
}
fn scope(&self, request_context: RequestContext) -> Cx {
Cx {
state: Arc::new(CxState {
shared: Arc::clone(&self.state.shared),
request_context: Arc::new(request_context),
tracker: self.state.tracker.clone(),
}),
identity: self.identity,
}
}
pub(crate) fn track(&self) -> (Cx, Arc<ContextTracker>) {
let tracker = Arc::new(ContextTracker::new(Arc::clone(&self.state.request_context)));
let child = Cx {
state: Arc::new(CxState {
shared: Arc::clone(&self.state.shared),
request_context: Arc::clone(&self.state.request_context),
tracker: Some(Arc::clone(&tracker)),
}),
identity: self.identity,
};
(child, tracker)
}
}
#[derive(Debug, Default)]
struct CxState {
shared: Arc<RequestShared>,
request_context: Arc<RequestContext>,
tracker: Option<Arc<ContextTracker>>,
}
#[derive(Debug, Default)]
struct RequestShared {
id: CxId,
app_context: Arc<AppContext>,
memoize_cache: MemoizeCache,
abort_store: AbortStore,
}
#[derive(Debug, Default)]
pub struct CxTestBuilder {
app_context: AppContext,
request_context: RequestContext,
}
impl CxTestBuilder {
#[must_use]
pub fn new() -> Self {
Self::default()
}
#[must_use]
pub fn app_context<T>(mut self, value: T) -> Self
where
T: Any + Send + Sync,
{
self.app_context.insert(value);
self
}
#[must_use]
pub fn request_context<T>(mut self, value: T) -> Self
where
T: Any + Send + Sync,
{
self.request_context.insert(value);
self
}
#[must_use]
pub fn build(self) -> Cx {
Cx::from_parts(Arc::new(self.app_context), self.request_context)
}
}
#[must_use]
#[track_caller]
pub fn identity(cx: &Cx) -> Identity {
match try_identity(cx) {
Ok(identity) => identity,
Err(error) => panic!("{error}"),
}
}
#[track_caller]
pub fn try_identity(cx: &Cx) -> Result<Identity, AmbiguousIdentityError> {
assert!(
cx.state.tracker.is_none(),
"identity cannot be read inside memoized functions"
);
cx.identity.checked()
}
#[doc(hidden)]
#[must_use]
pub fn identity_raw(cx: &Cx) -> Identity {
cx.identity
}
#[doc(hidden)]
#[must_use]
pub fn with_identity(cx: Cx, identity: Identity) -> Cx {
Cx { identity, ..cx }
}
#[inline]
#[must_use]
#[doc(hidden)]
pub fn memoize_cache(cx: &Cx) -> &MemoizeCache {
&cx.state.shared.memoize_cache
}
#[inline]
#[must_use]
#[doc(hidden)]
pub fn abort_store(cx: &Cx) -> &AbortStore {
&cx.state.shared.abort_store
}
#[cfg(test)]
mod tests {
use super::*;
#[derive(Debug, PartialEq)]
struct Marker(u32);
#[derive(Debug, PartialEq)]
struct Other(&'static str);
#[test]
fn a_fresh_context_has_a_unique_id() {
let first = Cx::new(Arc::new(AppContext::new()));
let second = Cx::new(Arc::new(AppContext::new()));
assert_ne!(first.id(), second.id());
}
#[test]
fn keys_distinguish_locations_and_repetitions() {
fn child(cx: &Cx, key: u32) -> Cx {
cx.keyed(key)
}
let cx = Cx::default();
assert_ne!(identity(&cx.keyed(())), identity(&cx.keyed(())));
assert_eq!(
identity(&child(&cx, 1)),
identity(&child(&Cx::default(), 1))
);
assert_ne!(identity(&child(&cx, 1)), identity(&child(&cx, 2)));
assert_ne!(
identity(&child(&child(&cx, 1), 2)),
identity(&child(&cx, 2))
);
assert_eq!(identity(&cx), Identity::ROOT);
}
#[test]
fn keyed_contexts_share_request_state_and_preserve_their_scope() {
let cx = Cx::default().with(Marker(7));
let child = cx.keyed("child");
assert!(Arc::ptr_eq(&child.state, &cx.state));
assert_eq!(child.id(), cx.id());
assert!(std::ptr::eq(memoize_cache(&child), memoize_cache(&cx)));
assert!(std::ptr::eq(
request_context::<Marker>(&child),
request_context::<Marker>(&cx)
));
let expected = identity(&child);
assert_eq!(identity(&child.clone()), expected);
assert_eq!(identity(&child.with(Other("value"))), expected);
assert_eq!(identity(&child.with_many((Other("value"),))), expected);
assert_eq!(identity_raw(&child.track().0), expected);
}
#[test]
#[should_panic(expected = "identity cannot be read inside memoized functions")]
fn memoized_functions_cannot_read_identity() {
let cx = Cx::default();
memoize_cache(&cx).memoize(&cx, (), (), |cx, ()| {
identity(&cx.keyed(()).with(Marker(7)))
});
}
#[test]
#[should_panic(expected = "identity cannot be read inside memoized functions")]
fn memoized_functions_cannot_try_identity() {
let cx = Cx::default();
memoize_cache(&cx).memoize(&cx, (), (), |cx, ()| try_identity(&cx));
}
#[tokio::test]
#[should_panic(expected = "identity cannot be read inside memoized functions")]
async fn memoized_functions_cannot_read_identity_after_suspension() {
let cx = Cx::default();
memoize_cache(&cx)
.memoize_async(&cx, (), (), |cx, ()| async move {
tokio::task::yield_now().await;
identity(&cx)
})
.await;
}
#[test]
fn with_registers_a_value_on_the_child() {
let cx = Cx::default();
let child = cx.with(Marker(1));
assert_eq!(try_request_context::<Marker>(&cx), None);
assert_eq!(request_context::<Marker>(&child), &Marker(1));
}
#[test]
fn with_shadows_without_touching_the_parent() {
let cx = Cx::default().with(Marker(1));
let child = cx.with(Marker(2));
assert_eq!(request_context::<Marker>(&cx), &Marker(1));
assert_eq!(request_context::<Marker>(&child), &Marker(2));
}
#[test]
fn a_child_inherits_the_parent_context() {
let cx = CxTestBuilder::new()
.app_context(Other("app"))
.request_context(Marker(7))
.build();
let child = cx.with(Other("request"));
assert_eq!(request_context::<Marker>(&child), &Marker(7));
assert_eq!(request_context::<Other>(&child), &Other("request"));
assert_eq!(app_context::<Other>(&child), &Other("app"));
}
#[test]
fn with_many_registers_every_value() {
let cx = Cx::default().with_many((Marker(1), Other("many")));
assert_eq!(request_context::<Marker>(&cx), &Marker(1));
assert_eq!(request_context::<Other>(&cx), &Other("many"));
}
#[test]
fn a_child_shares_the_request_state() {
let cx = Cx::default();
let child = cx.with(Marker(1));
assert_eq!(child.id(), cx.id());
assert!(std::ptr::eq(memoize_cache(&child), memoize_cache(&cx)));
assert!(std::ptr::eq(abort_store(&child), abort_store(&cx)));
}
#[test]
fn clones_outlive_the_original() {
let cx = CxTestBuilder::new().request_context(Marker(7)).build();
let id = cx.id();
let handle = cx.clone();
drop(cx);
assert_eq!(request_context::<Marker>(&handle).0, 7);
assert_eq!(handle.id(), id);
}
#[test]
fn handles_are_send_and_sync() {
fn assert_send_sync<T: Send + Sync>() {}
assert_send_sync::<Cx>();
}
}