use std::sync::Arc;
use crate::{Extract, InjectableResult, Provider, ResolveContext};
pub struct Inject<T: ?Sized>(pub Arc<T>);
impl<T: ?Sized> Clone for Inject<T> {
fn clone(&self) -> Self {
Inject(Arc::clone(&self.0))
}
}
impl<T: ?Sized + std::fmt::Debug> std::fmt::Debug for Inject<T> {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_tuple("Inject").field(&&*self.0).finish()
}
}
impl<T: ?Sized> Inject<T> {
pub fn new(value: Arc<T>) -> Self {
Self(value)
}
pub fn into_inner(self) -> Arc<T> {
self.0
}
pub fn inner(&self) -> &Arc<T> {
&self.0
}
pub fn arc(&self) -> Arc<T> {
Arc::clone(&self.0)
}
pub fn ptr_eq(&self, other: &Inject<T>) -> bool {
Arc::ptr_eq(&self.0, &other.0)
}
}
impl<T: ?Sized> std::ops::Deref for Inject<T> {
type Target = T;
fn deref(&self) -> &Self::Target {
&self.0
}
}
impl<T: ?Sized> From<Arc<T>> for Inject<T> {
fn from(arc: Arc<T>) -> Self {
Self(arc)
}
}
impl<T: ?Sized> From<Inject<T>> for Arc<T> {
fn from(inject: Inject<T>) -> Self {
inject.into_inner()
}
}
impl<T: PartialEq> PartialEq for Inject<T> {
fn eq(&self, other: &Self) -> bool {
**self == **other
}
}
impl<T: Eq> Eq for Inject<T> {}
impl<T: std::hash::Hash> std::hash::Hash for Inject<T> {
fn hash<H: std::hash::Hasher>(&self, state: &mut H) {
(**self).hash(state);
}
}
impl<T: ?Sized> AsRef<T> for Inject<T> {
fn as_ref(&self) -> &T {
self
}
}
impl<T: ?Sized> std::borrow::Borrow<T> for Inject<T> {
fn borrow(&self) -> &T {
self
}
}
#[async_trait::async_trait]
impl<T: Sized + Send + Sync + 'static> Extract for Inject<T> {
async fn extract(ctx: &ResolveContext) -> InjectableResult<Self> {
if let Some(result) = ctx.try_resolve_external::<Arc<T>>().await {
return result.map(Inject);
}
ctx.resolve_external::<T>()
.await
.map(|t| Inject(Arc::new(t)))
}
}
#[async_trait::async_trait]
impl<T: crate::Injectable> Extract for Arc<T> {
async fn extract(ctx: &ResolveContext) -> InjectableResult<Self> {
if T::IS_SINGLETON {
ctx.resolve_singleton_arc::<T>().await
} else {
let v = T::Provider::provide(ctx).await?;
Ok(Arc::new(v))
}
}
}
#[cfg(test)]
mod tests {
use super::*;
use std::collections::hash_map::DefaultHasher;
use std::hash::{Hash, Hasher};
fn make_inject(v: u32) -> Inject<u32> {
Inject::new(Arc::new(v))
}
#[test]
fn from_inject_into_arc() {
let inj = make_inject(42);
let arc: Arc<u32> = inj.into();
assert_eq!(*arc, 42);
}
#[test]
fn from_arc_into_inject() {
let arc = Arc::new(99u32);
let inj: Inject<u32> = arc.into();
assert_eq!(*inj, 99);
}
#[test]
fn partial_eq_same_value() {
let a = make_inject(1);
let b = make_inject(1);
assert_eq!(a, b);
}
#[test]
fn partial_eq_different_value() {
let a = make_inject(1);
let b = make_inject(2);
assert_ne!(a, b);
}
#[test]
fn hash_equals_inner_hash() {
let inj = make_inject(77);
let mut h1 = DefaultHasher::new();
inj.hash(&mut h1);
let mut h2 = DefaultHasher::new();
77u32.hash(&mut h2);
assert_eq!(h1.finish(), h2.finish());
}
#[test]
fn as_ref() {
let inj = make_inject(5);
let r: &u32 = inj.as_ref();
assert_eq!(*r, 5);
}
#[test]
fn borrow() {
use std::borrow::Borrow;
let inj = make_inject(10);
let b: &u32 = inj.borrow();
assert_eq!(*b, 10);
}
#[test]
fn debug_contains_inject() {
let inj = make_inject(7);
let s = format!("{inj:?}");
assert!(s.contains("Inject"));
assert!(s.contains('7'));
}
#[test]
fn clone_shares_arc() {
let inj = make_inject(3);
let cloned = inj.clone();
assert!(Arc::ptr_eq(&inj.0, &cloned.0));
}
#[test]
fn dyn_trait_inject_new() {
let arc: Arc<dyn std::fmt::Debug> = Arc::new(42u32);
let inj: Inject<dyn std::fmt::Debug> = Inject::new(arc);
let s = format!("{:?}", &*inj);
assert!(s.contains("42"));
}
}