1use std::sync::{Arc, Mutex};
13
14use futures::StreamExt;
15use serde::{Deserialize, Serialize};
16
17use crate::{
18 effect::{EffectId, EffectKind, HandlerDescriptor, Outcome, family},
19 error::{ErrorKind, ErrorReport},
20 wasm_compat::{WasmCompatSend, WasmCompatSync},
21};
22
23use super::{Dispatch, ErasedHandler, Reply, Serve, stream_truncated};
24
25#[derive(Debug, Clone, Serialize, Deserialize)]
27#[serde(tag = "decision", rename_all = "snake_case")]
28pub enum Decision {
29 Proceed,
31 Patch(EffectKind),
35 Deny(ErrorReport),
38}
39
40impl Decision {
41 pub fn deny(reason: impl Into<String>) -> Self {
43 Self::Deny(ErrorReport::new(ErrorKind::Denied, reason).with_retryable(false))
44 }
45}
46
47#[derive(Debug, Clone, Serialize, Deserialize)]
49#[serde(tag = "verdict", rename_all = "snake_case")]
50pub enum Verdict {
51 Keep,
53 Replace(Result<Outcome, ErrorReport>),
59}
60
61pub trait Intercept: WasmCompatSend + WasmCompatSync + 'static {
66 fn name(&self) -> String;
68
69 fn before(
71 &self,
72 id: EffectId,
73 kind: &EffectKind,
74 ) -> impl Future<Output = Decision> + WasmCompatSend;
75
76 fn after(
80 &self,
81 id: EffectId,
82 kind: &EffectKind,
83 outcome: &Result<Outcome, ErrorReport>,
84 ) -> impl Future<Output = Verdict> + WasmCompatSend;
85}
86
87pub struct Layer<I: Intercept> {
91 inner: ErasedHandler,
92 intercept: Arc<I>,
93}
94
95impl<I: Intercept> Layer<I> {
96 pub(crate) fn new(inner: ErasedHandler, intercept: I) -> Self {
98 Self {
99 inner,
100 intercept: Arc::new(intercept),
101 }
102 }
103
104 fn internal(&self, message: String) -> ErrorReport {
105 ErrorReport::new(
106 ErrorKind::Internal,
107 format!("layer `{}`: {message}", self.intercept.name()),
108 )
109 .with_retryable(false)
110 }
111}
112
113impl<I: Intercept> Serve for Layer<I> {
114 type Family = family::Dynamic;
115
116 fn descriptor(&self) -> HandlerDescriptor {
117 let mut descriptor = self.inner.descriptor();
118 descriptor.layers.insert(0, self.intercept.name());
119 descriptor
120 }
121
122 async fn serve(&self, kind: EffectKind, mut dispatch: Dispatch) -> Reply {
123 let id = dispatch.id();
124 let name = self.intercept.name();
125 let kind = match self.intercept.before(id, &kind).await {
126 Decision::Proceed => kind,
127 Decision::Patch(patched) => {
128 if patched.family() != kind.family() {
129 dispatch.discard(&name);
130 return Reply::Outcome(Err(self.internal(format!(
131 "patched a {} effect into a {} effect; a layer never changes the family",
132 kind.family(),
133 patched.family()
134 ))));
135 }
136 if let (
137 EffectKind::ToolCall { name: original, .. },
138 EffectKind::ToolCall {
139 name: replacement, ..
140 },
141 ) = (&kind, &patched)
142 && original != replacement
143 {
144 dispatch.discard(&name);
145 return Reply::Outcome(Err(self.internal(format!(
146 "patched tool target `{original}` into `{replacement}`; a layer never changes the bound tool"
147 ))));
148 }
149 dispatch.patched(&patched);
150 patched
151 }
152 Decision::Deny(report) => {
153 dispatch.discard(&name);
154 return Reply::Outcome(Err(report));
155 }
156 };
157 if !dispatch.is_stream() {
158 let folded = Arc::new(Mutex::new(None));
159 let inner = dispatch.inner(Some(folded.clone()));
160 let attribution = dispatch.attribution();
161 let outcome = self
162 .inner
163 .handle(kind.clone(), inner)
164 .await
165 .folded_outcome(Some(folded))
166 .await;
167 return Reply::Outcome(match self.intercept.after(id, &kind, &outcome).await {
168 Verdict::Keep => outcome,
169 Verdict::Replace(replacement) => {
170 attribution.replaced(&name);
171 replacement
172 }
173 });
174 }
175 let folded = Arc::new(Mutex::new(None));
176 let inner = dispatch.inner(Some(folded.clone()));
177 let attribution = dispatch.attribution();
178 let stream = self
179 .inner
180 .handle(kind.clone(), inner)
181 .await
182 .into_stream()
183 .fuse();
184 let intercept = self.intercept.clone();
185 Reply::Stream(Box::pin(futures::stream::unfold(
186 (stream, intercept, kind, folded, false, attribution),
187 move |(mut stream, intercept, kind, folded, mut decided, attribution)| async move {
188 let item = stream.next().await;
189 if decided && item.is_none() {
190 return None;
191 }
192 let outcome = if decided {
193 None
194 } else {
195 folded
196 .lock()
197 .unwrap_or_else(std::sync::PoisonError::into_inner)
198 .take()
199 };
200 let item = if let Some(outcome) = outcome {
201 decided = true;
202 match intercept.after(id, &kind, &outcome).await {
203 Verdict::Keep => item.unwrap_or_else(|| Err(stream_truncated())),
204 Verdict::Replace(Err(report)) => {
205 attribution.replaced(&intercept.name());
206 Err(report)
207 }
208 Verdict::Replace(Ok(_)) => {
209 attribution.replaced(&intercept.name());
210 Err(ErrorReport::new(
211 ErrorKind::Internal,
212 format!("layer `{}`: cannot replace a streamed answer already delivered; replace with an error, or decide before", intercept.name()),
213 ).with_retryable(false))
214 }
215 }
216 } else {
217 item?
218 };
219 Some((
220 item,
221 (stream, intercept, kind, folded, decided, attribution),
222 ))
223 },
224 )))
225 }
226}
227
228#[cfg(test)]
229mod tests;