1use 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
22pub 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
37pub 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 #[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 #[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 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 #[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 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 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 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 #[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}