Skip to main content

ruststream/runtime/
typed.rs

1//! Typed handler adapter: turns a handler over an input kind's borrowed target into a
2//! [`Handler<M>`](Handler) by materializing the input from each delivery via the input axis
3//! ([`InputKind`]).
4//!
5//! This is the decode boundary between the two middleware levels: raw (pre-decode) middleware
6//! wrap the produced `Handler<M>`; typed (post-decode) middleware wrap the `inner: Handler<T>`
7//! passed in here. Both use the same [`Layer`](super::Layer) / [`HandlerExt`](super::HandlerExt)
8//! machinery, just at different inputs.
9
10use std::{fmt, marker::PhantomData};
11
12use crate::IncomingMessage;
13use crate::codec::Codec;
14use serde::de::DeserializeOwned;
15use tracing::warn;
16
17use super::context::Context;
18use super::failure::FailurePolicy;
19use super::handler::{Handler, HandlerResult, Settle};
20use super::input::{DecodeWith, Decoded};
21
22/// Build a `Handler<M>` that decodes the payload with `codec` into `T` and forwards `&T` to
23/// `inner`.
24///
25/// `inner` is any [`Handler<T>`](Handler) - a closure `Fn(&T) -> _` or a typed middleware stack
26/// built with [`HandlerExt::with`](super::HandlerExt::with). The general form over any input
27/// kind (a raw `&[u8]` handler included) is [`Typed::over`].
28pub fn typed<M, T, C, H>(codec: C, inner: H) -> Typed<M, Decoded<T>, C, H>
29where
30    M: IncomingMessage,
31    T: DeserializeOwned + Send + Sync + 'static,
32    C: Codec,
33{
34    Typed::over(codec, inner)
35}
36
37/// Handler produced by [`typed`] (or [`Typed::over`] for a non-decoding input kind). Override
38/// the decode-failure policy with [`Typed::on_decode_failure`].
39pub struct Typed<M, Input, DecodeCodec, Inner> {
40    codec: DecodeCodec,
41    inner: Inner,
42    decode: FailurePolicy,
43    _phantom: PhantomData<fn(M, Input)>,
44}
45
46impl<M, Input, DecodeCodec, Inner> Typed<M, Input, DecodeCodec, Inner> {
47    /// Builds the adapter for any input kind: [`Decoded<T>`] decodes with `codec`,
48    /// [`RawBytes`](super::RawBytes) ignores it (pass `()`) and lends the payload itself.
49    #[must_use]
50    pub fn over(codec: DecodeCodec, inner: Inner) -> Self
51    where
52        Input: DecodeWith<DecodeCodec>,
53    {
54        Self {
55            codec,
56            inner,
57            decode: FailurePolicy::Drop,
58            _phantom: PhantomData,
59        }
60    }
61
62    /// Sets the [`FailurePolicy`] applied when the codec fails to decode an incoming payload. The
63    /// default is [`FailurePolicy::Drop`].
64    #[must_use]
65    pub fn on_decode_failure(mut self, decode: FailurePolicy) -> Self {
66        self.decode = decode;
67        self
68    }
69}
70
71impl<M, Input, DecodeCodec, Inner> fmt::Debug for Typed<M, Input, DecodeCodec, Inner> {
72    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
73        f.debug_struct("Typed")
74            .field("decode", &self.decode)
75            .finish_non_exhaustive()
76    }
77}
78
79impl<M, Input, DecodeCodec, Inner, Cx, St> Handler<M, Cx, St>
80    for Typed<M, Input, DecodeCodec, Inner>
81where
82    M: IncomingMessage,
83    Input: DecodeWith<DecodeCodec>,
84    DecodeCodec: Send + Sync,
85    Cx: Send,
86    St: Send + Sync,
87    Inner: Handler<Input::Target, Cx, St>,
88{
89    async fn handle(&self, msg: &M, ctx: &mut Context<'_, Cx, St>) -> Settle {
90        // The decode product lives on this stack frame and the handler borrows its view, so the
91        // input path allocates nothing of its own (a raw input borrows the payload straight out
92        // of the broker's buffer).
93        match Input::decode(&self.codec, msg.payload()) {
94            Ok(owned) => {
95                self.inner
96                    .handle(Input::view(&owned, msg.payload()), ctx)
97                    .await
98            }
99            Err(err) => {
100                warn!(
101                    target: "ruststream::dispatch",
102                    subscription = %ctx.name(),
103                    message_type = Input::input_label(),
104                    error = %err,
105                    "codec decode failed",
106                );
107                #[cfg(any(feature = "testing", feature = "otel"))]
108                ctx.mark_decode_failed();
109                match self.decode {
110                    FailurePolicy::FailFast => {
111                        ctx.fail_fast(&format!("decode failed: {err}"));
112                        HandlerResult::drop()
113                    }
114                    other => other.settlement().unwrap_or_else(HandlerResult::drop),
115                }
116                .into()
117            }
118        }
119    }
120}
121
122#[cfg(all(test, feature = "json"))]
123mod tests {
124    use std::sync::{
125        Arc,
126        atomic::{AtomicU32, Ordering},
127    };
128
129    use super::typed;
130    use crate::codec::JsonCodec;
131    use crate::runtime::context::Context;
132    use crate::runtime::dispatch::Delivery;
133    use crate::runtime::failure::FailurePolicy;
134    use crate::runtime::handler::{Handler, HandlerResult};
135    use crate::{AckError, Headers, IncomingMessage};
136
137    struct StubMsg(Vec<u8>, Headers);
138
139    impl IncomingMessage for StubMsg {
140        fn payload(&self) -> &[u8] {
141            &self.0
142        }
143
144        fn headers(&self) -> &Headers {
145            &self.1
146        }
147
148        async fn ack(self) -> Result<(), AckError> {
149            Ok(())
150        }
151
152        async fn nack(self, _requeue: bool) -> Result<(), AckError> {
153            Ok(())
154        }
155    }
156
157    fn counting_inner(seen: &Arc<AtomicU32>) -> impl Handler<u32> {
158        let seen = Arc::clone(seen);
159        move |value: &u32, _ctx: &mut Context| {
160            let seen = Arc::clone(&seen);
161            let value = *value;
162            async move {
163                seen.store(value, Ordering::SeqCst);
164                HandlerResult::Ack
165            }
166        }
167    }
168
169    // Plain #[tokio::test]: nothing is spawned, the handler future is awaited inline.
170    #[tokio::test]
171    async fn decoded_value_reaches_inner() {
172        let seen = Arc::new(AtomicU32::new(0));
173        let handler = typed(JsonCodec, counting_inner(&seen));
174        let state = ();
175        let delivery = Delivery::empty();
176        let headers = Headers::new();
177        let mut ctx = Context::new("typed", &headers, &state, (), &delivery);
178
179        let msg = StubMsg(b"7".to_vec(), Headers::new());
180        assert_eq!(
181            handler.handle(&msg, &mut ctx).await.outcome(),
182            HandlerResult::Ack
183        );
184        assert_eq!(seen.load(Ordering::SeqCst), 7);
185    }
186
187    #[tokio::test]
188    async fn raw_bytes_lend_the_payload_itself() {
189        use super::Typed;
190        use crate::runtime::input::RawBytes;
191
192        let seen = Arc::new(AtomicU32::new(0));
193        let inner = {
194            let seen = Arc::clone(&seen);
195            move |bytes: &[u8], _ctx: &mut Context| {
196                let seen = Arc::clone(&seen);
197                let len = u32::try_from(bytes.len()).unwrap();
198                async move {
199                    seen.store(len, Ordering::SeqCst);
200                    HandlerResult::Ack
201                }
202            }
203        };
204        // No codec anywhere: the raw kind decodes with `()`.
205        let handler = Typed::<StubMsg, RawBytes, (), _>::over((), inner);
206        let state = ();
207        let delivery = Delivery::empty();
208        let headers = Headers::new();
209        let mut ctx = Context::new("frames", &headers, &state, (), &delivery);
210
211        let msg = StubMsg(b"not json at all".to_vec(), Headers::new());
212        assert_eq!(
213            handler.handle(&msg, &mut ctx).await.outcome(),
214            HandlerResult::Ack
215        );
216        assert_eq!(seen.load(Ordering::SeqCst), 15);
217    }
218
219    #[tokio::test]
220    async fn decode_failure_drops_by_default() {
221        let seen = Arc::new(AtomicU32::new(0));
222        let handler = typed(JsonCodec, counting_inner(&seen));
223        let state = ();
224        let delivery = Delivery::empty();
225        let headers = Headers::new();
226        let mut ctx = Context::new("typed", &headers, &state, (), &delivery);
227
228        let msg = StubMsg(b"not json".to_vec(), Headers::new());
229        assert_eq!(
230            handler.handle(&msg, &mut ctx).await.outcome(),
231            HandlerResult::drop()
232        );
233        assert_eq!(seen.load(Ordering::SeqCst), 0, "inner must not run");
234    }
235
236    #[tokio::test]
237    async fn decode_failure_requeues_when_overridden() {
238        let seen = Arc::new(AtomicU32::new(0));
239        let handler =
240            typed(JsonCodec, counting_inner(&seen)).on_decode_failure(FailurePolicy::Retry);
241        let state = ();
242        let delivery = Delivery::empty();
243        let headers = Headers::new();
244        let mut ctx = Context::new("typed", &headers, &state, (), &delivery);
245
246        let msg = StubMsg(b"not json".to_vec(), Headers::new());
247        assert_eq!(
248            handler.handle(&msg, &mut ctx).await.outcome(),
249            HandlerResult::retry()
250        );
251        assert_eq!(seen.load(Ordering::SeqCst), 0, "inner must not run");
252    }
253
254    #[tokio::test]
255    async fn typed_handler_is_debug_and_stub_acks() {
256        let seen = Arc::new(AtomicU32::new(0));
257        let handler = typed(JsonCodec, counting_inner(&seen));
258        let state = ();
259        let delivery = Delivery::empty();
260        let headers = Headers::new();
261        let mut ctx = Context::new("typed", &headers, &state, (), &delivery);
262        // Drive one delivery to pin the message type, then check the Debug rendering.
263        let msg = StubMsg(b"5".to_vec(), Headers::new());
264        let _ = handler.handle(&msg, &mut ctx).await;
265        assert!(format!("{handler:?}").contains("Typed"));
266
267        // Exercise the StubMsg fixture's own IncomingMessage surface.
268        let other = StubMsg(b"x".to_vec(), Headers::new());
269        assert!(other.headers().is_empty());
270        other.ack().await.unwrap();
271        StubMsg(Vec::new(), Headers::new())
272            .nack(true)
273            .await
274            .unwrap();
275    }
276
277    // Captures the fields of the one event emitted on a decode failure, so the test can assert the
278    // diagnostic carries the subscription name and target type (needs a tracing subscriber, hence
279    // the `logging` feature gate).
280    #[cfg(feature = "logging")]
281    #[tokio::test]
282    async fn decode_failure_log_names_subscription_and_type() {
283        use std::collections::HashMap;
284        use std::sync::Mutex;
285
286        use tracing::field::{Field, Visit};
287        use tracing_subscriber::Layer;
288        use tracing_subscriber::layer::{Context as LayerContext, SubscriberExt as _};
289
290        #[derive(Default)]
291        struct FieldGrab(HashMap<String, String>);
292
293        impl Visit for FieldGrab {
294            fn record_str(&mut self, field: &Field, value: &str) {
295                self.0.insert(field.name().to_owned(), value.to_owned());
296            }
297
298            fn record_debug(&mut self, field: &Field, value: &dyn std::fmt::Debug) {
299                self.0
300                    .entry(field.name().to_owned())
301                    .or_insert_with(|| format!("{value:?}"));
302            }
303        }
304
305        struct Capture(Arc<Mutex<Vec<HashMap<String, String>>>>);
306
307        impl<S: tracing::Subscriber> Layer<S> for Capture {
308            fn on_event(&self, event: &tracing::Event<'_>, _ctx: LayerContext<'_, S>) {
309                let mut grab = FieldGrab::default();
310                event.record(&mut grab);
311                self.0.lock().unwrap().push(grab.0);
312            }
313        }
314
315        let events = Arc::new(Mutex::new(Vec::new()));
316        let guard = tracing::subscriber::set_default(
317            tracing_subscriber::registry().with(Capture(Arc::clone(&events))),
318        );
319
320        let seen = Arc::new(AtomicU32::new(0));
321        let handler = typed(JsonCodec, counting_inner(&seen));
322        let state = ();
323        let delivery = Delivery::empty();
324        let headers = Headers::new();
325        let mut ctx = Context::new("orders.inbound", &headers, &state, (), &delivery);
326        let msg = StubMsg(b"not json".to_vec(), Headers::new());
327        assert_eq!(
328            handler.handle(&msg, &mut ctx).await.outcome(),
329            HandlerResult::drop()
330        );
331        drop(guard);
332
333        let decode_event = {
334            let captured = events.lock().unwrap();
335            captured
336                .iter()
337                .find(|f| f.get("message").is_some_and(|m| m == "codec decode failed"))
338                .cloned()
339                .expect("a codec-decode-failed event must be emitted")
340        };
341        assert_eq!(
342            decode_event.get("subscription").map(String::as_str),
343            Some("orders.inbound")
344        );
345        assert_eq!(
346            decode_event.get("message_type").map(String::as_str),
347            Some("u32")
348        );
349    }
350}