use std::{
sync::{Arc, Mutex},
task::Poll,
};
use futures::{StreamExt, channel::oneshot};
use crate::{
effect::{EffectId, EffectKind, HandlerDescriptor, Outcome},
error::{ErrorKind, ErrorReport},
streaming::{Item, Relayed, StreamEvent, StreamEvents},
wasm_compat::{WasmBoxedFuture, WasmCompatSend, WasmCompatSync},
};
#[cfg(test)]
mod tests;
pub type HandlerFuture<'a> = WasmBoxedFuture<'a, Reply>;
pub trait Serve: WasmCompatSend + WasmCompatSync {
type Family: crate::effect::Served;
fn descriptor(&self) -> HandlerDescriptor;
fn serve(
&self,
kind: EffectKind,
dispatch: Dispatch,
) -> impl Future<Output = Reply> + WasmCompatSend + use<'_, Self>;
}
pub(crate) trait Handler: WasmCompatSend + WasmCompatSync {
fn descriptor(&self) -> HandlerDescriptor;
fn handle(&self, kind: EffectKind, dispatch: Dispatch) -> HandlerFuture<'_>;
}
#[diagnostic::do_not_recommend]
impl<T: Serve> Handler for T {
fn descriptor(&self) -> HandlerDescriptor {
Serve::descriptor(self)
}
fn handle(&self, kind: EffectKind, dispatch: Dispatch) -> HandlerFuture<'_> {
let observer = dispatch.observer.clone();
let folded = dispatch.folded.clone();
let streaming = dispatch.is_stream();
Box::pin(async move {
let reply = self.serve(kind, dispatch).await;
let seen = observer.and_then(|slot| lock(&slot).take());
reply.observed(streaming, seen, folded)
})
}
}
impl<H: Serve + ?Sized> Serve for Arc<H> {
type Family = H::Family;
fn descriptor(&self) -> HandlerDescriptor {
(**self).descriptor()
}
async fn serve(&self, kind: EffectKind, dispatch: Dispatch) -> Reply {
(**self).serve(kind, dispatch).await
}
}
#[derive(Clone)]
pub struct ErasedHandler(ErasedInner);
#[cfg(not(target_family = "wasm"))]
type ErasedInner = Arc<dyn Handler + Send + Sync>;
#[cfg(target_family = "wasm")]
type ErasedInner = Arc<dyn Handler>;
impl ErasedHandler {
pub fn new(handler: impl Serve + 'static) -> Self {
Self(Arc::new(handler))
}
pub fn layered(self, intercept: impl super::Intercept) -> Self {
Self::new(super::Layer::new(self, intercept))
}
pub fn descriptor(&self) -> HandlerDescriptor {
self.0.descriptor()
}
pub fn handle(&self, kind: EffectKind, dispatch: Dispatch) -> HandlerFuture<'_> {
self.0.handle(kind, dispatch)
}
pub fn ptr_eq(&self, other: &Self) -> bool {
Arc::ptr_eq(&self.0, &other.0)
}
}
impl std::fmt::Debug for ErasedHandler {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("ErasedHandler")
.field("key", &self.0.descriptor().key)
.finish_non_exhaustive()
}
}
impl Serve for ErasedHandler {
type Family = crate::effect::family::Dynamic;
fn descriptor(&self) -> HandlerDescriptor {
self.0.descriptor()
}
async fn serve(&self, kind: EffectKind, dispatch: Dispatch) -> Reply {
self.handle(kind, dispatch).await
}
}
pub trait Observe: Send + Sync {
fn adapter_context(&self) -> Option<crate::observe::AdapterContext> {
None
}
fn outcome(&mut self, outcome: &Result<Outcome, ErrorReport>);
fn keep_events(&self) -> bool;
fn event(&mut self, item: &Item<StreamEvent>);
fn stream_error(&mut self, _error: &ErrorReport) {}
fn origin(&mut self, origin: &crate::message::Origin);
fn stream_item(
&mut self,
item: &Result<Relayed, ErrorReport>,
outcome: Option<&Result<Outcome, ErrorReport>>,
) {
if self.keep_events() {
match item {
Ok(Relayed::Origin(origin)) => self.origin(origin),
Ok(Relayed::Item(item)) => self.event(item),
Ok(Relayed::Done(_)) => {}
Err(error) => self.stream_error(error),
}
}
if let Some(outcome) = outcome {
self.outcome(outcome);
}
}
fn discard(&mut self, layer: &str);
fn patch(&mut self, kind: &EffectKind);
}
#[derive(Default)]
pub struct StreamTap {
finished: bool,
}
impl StreamTap {
pub fn new() -> Self {
Self::default()
}
pub fn observe(
&mut self,
item: &Result<Relayed, ErrorReport>,
) -> Option<Result<Outcome, ErrorReport>> {
if self.finished {
return None;
}
let outcome = match item {
Ok(Relayed::Origin(_) | Relayed::Item(_)) => None,
Ok(Relayed::Done(response)) => Some(Ok(Outcome::Completion((**response).clone()))),
Err(report) => Some(Err(report.clone())),
};
self.finished = outcome.is_some();
outcome
}
}
pub fn stream_truncated() -> ErrorReport {
ErrorReport::from(&crate::error::ProviderError::Truncated)
}
pub enum Reply {
Outcome(Result<Outcome, ErrorReport>),
Stream(StreamEvents),
}
impl Reply {
pub async fn into_outcome(self) -> Result<Outcome, ErrorReport> {
self.folded_outcome(None).await
}
pub(crate) async fn folded_outcome(
self,
folded: Option<Folded>,
) -> Result<Outcome, ErrorReport> {
match self {
Self::Outcome(outcome) => outcome,
Self::Stream(mut stream) => {
let mut fold = StreamTap::new();
while let Some(item) = stream.next().await {
let outcome = match &folded {
Some(folded) => lock(folded).take(),
None => fold.observe(&item),
};
if let Some(outcome) = outcome {
return outcome;
}
}
Err(stream_truncated())
}
}
}
pub fn into_stream(self) -> StreamEvents {
match self {
Self::Stream(stream) => stream,
Self::Outcome(outcome) => Box::pin(futures::stream::iter(match outcome {
Ok(Outcome::Completion(response)) => {
match crate::operation::completion::events_of(&response) {
Ok(items) => std::iter::once(Ok(Relayed::Origin(response.origin.clone())))
.chain(items.into_iter().map(|item| Ok(Relayed::Item(item))))
.chain(std::iter::once(Ok(Relayed::Done(Box::new(response)))))
.collect(),
Err(error) => vec![Err(ErrorReport::from(&error))],
}
}
Ok(other) => vec![Err(wrong_stream_answer(&other))],
Err(report) => vec![Err(report)],
})),
}
}
fn observed(self, streaming: bool, mut seen: Option<Observed>, folded: Option<Folded>) -> Self {
if !streaming && let Self::Outcome(outcome) = self {
if let Some(seen) = &mut seen {
seen.outcome(&outcome);
}
return Self::Outcome(outcome);
}
if seen.is_none() && folded.is_none() {
return if streaming {
Self::Stream(self.into_stream())
} else {
self
};
}
let original = match &self {
Self::Outcome(Ok(Outcome::Completion(response))) => {
Some(Ok(Outcome::Completion(response.clone())))
}
_ => None,
};
let mut stream = self.into_stream();
let mut fold = StreamTap::new();
let mut finished = false;
Self::Stream(Box::pin(futures::stream::poll_fn(move |cx| {
let item = match stream.as_mut().poll_next(cx) {
Poll::Pending => return Poll::Pending,
Poll::Ready(item) => item,
};
if let Some(item) = &item {
let outcome = if finished { None } else { fold.observe(item) };
if let Some(seen) = &mut seen {
let recorded = outcome.as_ref().map(|outcome| {
if matches!(item, Ok(Relayed::Done(_))) {
original.as_ref().unwrap_or(outcome)
} else {
outcome
}
});
if streaming {
seen.item(item, recorded);
} else if let Some(recorded) = recorded {
seen.outcome(recorded);
}
}
if let Some(outcome) = outcome {
finished = true;
if let Some(folded) = &folded {
*lock(folded) = Some(outcome);
}
}
} else if !finished {
finished = true;
let outcome = Err(stream_truncated());
if let Some(seen) = &mut seen {
seen.outcome(&outcome);
}
if let Some(folded) = &folded {
*lock(folded) = Some(outcome);
}
}
Poll::Ready(item)
})))
}
}
pub(crate) type Folded = Arc<Mutex<Option<Result<Outcome, ErrorReport>>>>;
fn lock<T>(value: &Mutex<T>) -> std::sync::MutexGuard<'_, T> {
value
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner)
}
pub struct Dispatch {
adapter_context: Option<crate::observe::AdapterContext>,
adapter_context_explicit: bool,
id: EffectId,
streaming: bool,
scopes: Vec<Arc<dyn std::any::Any + Send + Sync>>,
observer: Option<Arc<Mutex<Option<Observed>>>>,
replaced_by: Arc<Mutex<Option<String>>>,
folded: Option<Folded>,
}
#[derive(Clone)]
pub(crate) struct Attribution(Arc<Mutex<Option<String>>>);
impl Attribution {
pub(crate) fn replaced(&self, layer: &str) {
*lock(&self.0) = Some(layer.to_owned());
}
}
impl Dispatch {
pub fn new(id: EffectId, streaming: bool) -> Self {
Self {
id,
streaming,
adapter_context: None,
adapter_context_explicit: false,
scopes: Vec::new(),
observer: None,
replaced_by: Arc::new(Mutex::new(None)),
folded: None,
}
}
pub fn replaced_by(&self) -> Arc<Mutex<Option<String>>> {
self.replaced_by.clone()
}
pub const fn id(&self) -> EffectId {
self.id
}
pub const fn is_stream(&self) -> bool {
self.streaming
}
pub fn with_scope(mut self, scope: Arc<dyn std::any::Any + Send + Sync>) -> Self {
self.scopes.push(scope);
self
}
pub fn scope<T: std::any::Any + Send + Sync>(&self) -> Option<Arc<T>> {
self.scopes
.iter()
.find_map(|scope| Arc::downcast::<T>(scope.clone()).ok())
}
pub fn scopes(&self) -> Vec<Arc<dyn std::any::Any + Send + Sync>> {
self.scopes.clone()
}
pub fn with_observer(mut self, observer: Box<dyn Observe>) -> Self {
if !self.adapter_context_explicit {
self.adapter_context = observer.adapter_context();
}
self.observer = Some(Arc::new(Mutex::new(Some(Observed {
observer,
told: false,
}))));
self
}
pub fn with_adapter_context(mut self, context: crate::observe::AdapterContext) -> Self {
self.adapter_context = Some(context);
self.adapter_context_explicit = true;
self
}
pub fn adapter_context(&self) -> Option<crate::observe::AdapterContext> {
self.adapter_context.clone()
}
pub(crate) fn patched(&mut self, kind: &EffectKind) {
if let Some(slot) = &self.observer
&& let Some(seen) = lock(slot).as_mut()
{
seen.observer.patch(kind);
}
}
pub(crate) fn discard(&mut self, layer: &str) {
if let Some(slot) = &self.observer
&& let Some(mut seen) = lock(slot).take()
{
seen.told = true;
seen.observer.discard(layer);
}
}
pub(crate) fn attribution(&self) -> Attribution {
Attribution(self.replaced_by.clone())
}
pub(crate) fn inner(&mut self, folded: Option<Folded>) -> Self {
let observer = self.observer.as_ref().and_then(|slot| lock(slot).take());
Self {
id: self.id,
streaming: self.streaming,
adapter_context: self.adapter_context.clone(),
adapter_context_explicit: self.adapter_context_explicit,
scopes: self.scopes.clone(),
observer: observer.map(|seen| Arc::new(Mutex::new(Some(seen)))),
replaced_by: self.replaced_by.clone(),
folded,
}
}
}
struct Observed {
observer: Box<dyn Observe>,
told: bool,
}
impl Observed {
fn outcome(&mut self, outcome: &Result<Outcome, ErrorReport>) {
if !self.told {
self.told = true;
self.observer.outcome(outcome);
}
}
fn item(
&mut self,
item: &Result<Relayed, ErrorReport>,
outcome: Option<&Result<Outcome, ErrorReport>>,
) {
let outcome = outcome.filter(|_| !self.told);
self.told |= outcome.is_some();
self.observer.stream_item(item, outcome);
}
}
impl Drop for Observed {
fn drop(&mut self) {
self.outcome(&Err(cancelled()));
}
}
fn wrong_stream_answer(other: &Outcome) -> ErrorReport {
ErrorReport::new(
ErrorKind::Internal,
format!(
"a streaming dispatch was answered with a {} outcome",
other.family()
),
)
}
pub fn cancelled() -> ErrorReport {
ErrorReport::new(
ErrorKind::Cancelled,
"the consumer cancelled the dispatch before it was answered",
)
.with_retryable(false)
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, thiserror::Error)]
#[error("the dispatch's consumer is gone")]
pub struct SinkClosed;
pub struct Resolver(oneshot::Sender<Result<Outcome, ErrorReport>>);
pub fn deferred() -> (
Resolver,
impl Future<Output = Result<Outcome, ErrorReport>> + Send + 'static,
) {
let (sender, receiver) = oneshot::channel();
(Resolver(sender), async move {
receiver.await.unwrap_or_else(|_| {
Err(ErrorReport::new(
ErrorKind::Internal,
"the handler dropped its outcome sink without answering",
))
})
})
}
impl Resolver {
pub fn resolve(self, outcome: Result<Outcome, ErrorReport>) -> Result<(), SinkClosed> {
self.0.send(outcome).map_err(|_| SinkClosed)
}
pub fn is_closed(&self) -> bool {
self.0.is_canceled()
}
}
pub async fn serve_inline(
handler: &ErasedHandler,
kind: EffectKind,
) -> Result<Outcome, ErrorReport> {
serve_inline_with(handler, kind, Vec::new()).await
}
pub async fn serve_inline_with(
handler: &ErasedHandler,
kind: EffectKind,
scopes: Vec<Arc<dyn std::any::Any + Send + Sync>>,
) -> Result<Outcome, ErrorReport> {
let mut dispatch = Dispatch::new(EffectId::from_raw(0), false);
for scope in scopes {
dispatch = dispatch.with_scope(scope);
}
handler.handle(kind, dispatch).await.into_outcome().await
}