Skip to main content

rig_core/serve/
layer.rs

1//! Handler composition with pre-dispatch decisions and post-answer verdicts.
2//! Recording observes the innermost handler: denied dispatches leave no record,
3//! and replaced answers retain the handler's original outcome for replay.
4//!
5//! ```
6//! use rig_core::serve::Decision;
7//!
8//! let decision = Decision::deny("Not authorized");
9//! assert!(matches!(decision, Decision::Deny(_)));
10//! ```
11
12use 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/// What a layer decides about a dispatch before the handler sees it.
26#[derive(Debug, Clone, Serialize, Deserialize)]
27#[serde(tag = "decision", rename_all = "snake_case")]
28pub enum Decision {
29    /// Serve it as it is.
30    Proceed,
31    /// Serve this instead. A patch never changes the family or a tool call's
32    /// target name: either change returns an `Internal` report without reaching
33    /// the next layer or dispatching. Tool argument patches remain supported.
34    Patch(EffectKind),
35    /// Returns this report without invoking the handler or recording a dispatch.
36    /// The report's kind is preserved; [`Decision::deny`] uses `ErrorKind::Denied`.
37    Deny(ErrorReport),
38}
39
40impl Decision {
41    /// A denial by policy: `ErrorKind::Denied`, never retryable.
42    pub fn deny(reason: impl Into<String>) -> Self {
43        Self::Deny(ErrorReport::new(ErrorKind::Denied, reason).with_retryable(false))
44    }
45}
46
47/// What a layer decides about an answer on its way out.
48#[derive(Debug, Clone, Serialize, Deserialize)]
49#[serde(tag = "verdict", rename_all = "snake_case")]
50pub enum Verdict {
51    /// The consumer receives what the handler answered.
52    Keep,
53    /// The consumer receives this instead; the record keeps the handler's
54    /// answer. Over a streaming dispatch the events were already delivered
55    /// as they came, so only an error can replace the answer there: a
56    /// `Replace(Ok(_))` reaches the consumer as an `Internal` error naming
57    /// the layer.
58    Replace(Result<Outcome, ErrorReport>),
59}
60
61/// Host policy before a handler and after its first answer. A suspended
62/// verdict keeps execution in flight; driver cancellation drops its future.
63/// Streamed verdict futures are polled on the stream consumer's thread.
64/// The layer name identifies policy when recording and validating replay.
65pub trait Intercept: WasmCompatSend + WasmCompatSync + 'static {
66    /// The layer's name, as the log records it.
67    fn name(&self) -> String;
68
69    /// Before the handler: the dispatch as it will be served, or not.
70    fn before(
71        &self,
72        id: EffectId,
73        kind: &EffectKind,
74    ) -> impl Future<Output = Decision> + WasmCompatSend;
75
76    /// After the handler: the answer as the consumer will receive it. For a
77    /// streaming dispatch `outcome` is the fold of the events (the
78    /// completion the record stores).
79    fn after(
80        &self,
81        id: EffectId,
82        kind: &EffectKind,
83        outcome: &Result<Outcome, ErrorReport>,
84    ) -> impl Future<Output = Verdict> + WasmCompatSend;
85}
86
87/// A handler wrapped in a policy: a [`Serve`] like any other, registered
88/// under the inner handler's descriptor (with the layer's name added,
89/// outermost first). Built with [`ErasedHandler::layered`].
90pub struct Layer<I: Intercept> {
91    inner: ErasedHandler,
92    intercept: Arc<I>,
93}
94
95impl<I: Intercept> Layer<I> {
96    /// `intercept` around `inner`.
97    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;