use std::sync::{Arc, Mutex};
use futures::StreamExt;
use serde::{Deserialize, Serialize};
use crate::{
effect::{EffectId, EffectKind, HandlerDescriptor, Outcome, family},
error::{ErrorKind, ErrorReport},
wasm_compat::{WasmCompatSend, WasmCompatSync},
};
use super::{Dispatch, ErasedHandler, Reply, Serve, stream_truncated};
#[derive(Debug, Clone, Serialize, Deserialize)]
#[serde(tag = "decision", rename_all = "snake_case")]
pub enum Decision {
Proceed,
Patch(EffectKind),
Deny(ErrorReport),
}
impl Decision {
pub fn deny(reason: impl Into<String>) -> Self {
Self::Deny(ErrorReport::new(ErrorKind::Denied, reason).with_retryable(false))
}
}
#[derive(Debug, Clone, Serialize, Deserialize)]
#[serde(tag = "verdict", rename_all = "snake_case")]
pub enum Verdict {
Keep,
Replace(Result<Outcome, ErrorReport>),
}
pub trait Intercept: WasmCompatSend + WasmCompatSync + 'static {
fn name(&self) -> String;
fn before(
&self,
id: EffectId,
kind: &EffectKind,
) -> impl Future<Output = Decision> + WasmCompatSend;
fn after(
&self,
id: EffectId,
kind: &EffectKind,
outcome: &Result<Outcome, ErrorReport>,
) -> impl Future<Output = Verdict> + WasmCompatSend;
}
pub struct Layer<I: Intercept> {
inner: ErasedHandler,
intercept: Arc<I>,
}
impl<I: Intercept> Layer<I> {
pub(crate) fn new(inner: ErasedHandler, intercept: I) -> Self {
Self {
inner,
intercept: Arc::new(intercept),
}
}
fn internal(&self, message: String) -> ErrorReport {
ErrorReport::new(
ErrorKind::Internal,
format!("layer `{}`: {message}", self.intercept.name()),
)
.with_retryable(false)
}
}
impl<I: Intercept> Serve for Layer<I> {
type Family = family::Dynamic;
fn descriptor(&self) -> HandlerDescriptor {
let mut descriptor = self.inner.descriptor();
descriptor.layers.insert(0, self.intercept.name());
descriptor
}
async fn serve(&self, kind: EffectKind, mut dispatch: Dispatch) -> Reply {
let id = dispatch.id();
let name = self.intercept.name();
let kind = match self.intercept.before(id, &kind).await {
Decision::Proceed => kind,
Decision::Patch(patched) => {
if patched.family() != kind.family() {
dispatch.discard(&name);
return Reply::Outcome(Err(self.internal(format!(
"patched a {} effect into a {} effect; a layer never changes the family",
kind.family(),
patched.family()
))));
}
if let (
EffectKind::ToolCall { name: original, .. },
EffectKind::ToolCall {
name: replacement, ..
},
) = (&kind, &patched)
&& original != replacement
{
dispatch.discard(&name);
return Reply::Outcome(Err(self.internal(format!(
"patched tool target `{original}` into `{replacement}`; a layer never changes the bound tool"
))));
}
dispatch.patched(&patched);
patched
}
Decision::Deny(report) => {
dispatch.discard(&name);
return Reply::Outcome(Err(report));
}
};
if !dispatch.is_stream() {
let folded = Arc::new(Mutex::new(None));
let inner = dispatch.inner(Some(folded.clone()));
let attribution = dispatch.attribution();
let outcome = self
.inner
.handle(kind.clone(), inner)
.await
.folded_outcome(Some(folded))
.await;
return Reply::Outcome(match self.intercept.after(id, &kind, &outcome).await {
Verdict::Keep => outcome,
Verdict::Replace(replacement) => {
attribution.replaced(&name);
replacement
}
});
}
let folded = Arc::new(Mutex::new(None));
let inner = dispatch.inner(Some(folded.clone()));
let attribution = dispatch.attribution();
let stream = self
.inner
.handle(kind.clone(), inner)
.await
.into_stream()
.fuse();
let intercept = self.intercept.clone();
Reply::Stream(Box::pin(futures::stream::unfold(
(stream, intercept, kind, folded, false, attribution),
move |(mut stream, intercept, kind, folded, mut decided, attribution)| async move {
let item = stream.next().await;
if decided && item.is_none() {
return None;
}
let outcome = if decided {
None
} else {
folded
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner)
.take()
};
let item = if let Some(outcome) = outcome {
decided = true;
match intercept.after(id, &kind, &outcome).await {
Verdict::Keep => item.unwrap_or_else(|| Err(stream_truncated())),
Verdict::Replace(Err(report)) => {
attribution.replaced(&intercept.name());
Err(report)
}
Verdict::Replace(Ok(_)) => {
attribution.replaced(&intercept.name());
Err(ErrorReport::new(
ErrorKind::Internal,
format!("layer `{}`: cannot replace a streamed answer already delivered; replace with an error, or decide before", intercept.name()),
).with_retryable(false))
}
}
} else {
item?
};
Some((
item,
(stream, intercept, kind, folded, decided, attribution),
))
},
)))
}
}
#[cfg(test)]
mod tests;