1use std::marker::PhantomData;
2use std::pin::Pin;
3
4use futures::{StreamExt, stream};
5
6use crate::error::Result;
7use crate::io::message::Message;
8use crate::io::{Acker, BoxStream, Reader};
9use crate::payload::Payload;
10use crate::payload_codec::{EventCodec, PayloadCodec, PayloadEventCodec};
11
12pub struct DecodeReader<R, C, P> {
15 inner: R,
16 codec: C,
17 disposition: DecodeErrorDisposition,
18 _payload: PhantomData<P>,
19}
20
21#[derive(Debug, Clone, Copy, PartialEq, Eq, Default)]
22pub enum DecodeErrorDisposition {
23 Surface,
24 #[default]
25 AckInner,
26 NackInner,
27}
28
29impl<R, C, P> DecodeReader<R, C, P> {
30 pub fn new(inner: R, codec: C) -> Self {
31 Self {
32 inner,
33 codec,
34 disposition: DecodeErrorDisposition::default(),
35 _payload: PhantomData,
36 }
37 }
38
39 pub fn with_disposition(mut self, disposition: DecodeErrorDisposition) -> Self {
40 self.disposition = disposition;
41 self
42 }
43}
44
45impl<R, C, P> DecodeReader<R, PayloadEventCodec<C>, P>
46where
47 C: PayloadCodec<P>,
48{
49 pub fn from_payload_codec(inner: R, codec: C) -> Self {
50 Self::new(inner, PayloadEventCodec::new(codec))
51 }
52}
53
54impl<R, C, P> Reader<P> for DecodeReader<R, C, P>
55where
56 R: Reader<Payload> + Send + Sync + 'static,
57 R::Subscription: Send + 'static,
58 R::Acker: Acker + Send + Sync + 'static,
59 R::Cursor: Send + Sync + 'static,
60 R::Stream: 'static,
61 C: EventCodec<P> + Clone + Send + Sync + 'static,
62 P: Send + Sync + 'static,
63{
64 type Subscription = R::Subscription;
65 type Acker = R::Acker;
66 type Cursor = R::Cursor;
67 type Stream = BoxStream<R::Cursor, R::Acker, P>;
68
69 async fn read(&self, subscription: Self::Subscription) -> Result<Self::Stream> {
70 let inner = self.inner.read(subscription).await?;
71 let state = DecodeState {
72 inner: Box::pin(inner),
73 codec: self.codec.clone(),
74 disposition: self.disposition,
75 _payload: PhantomData,
76 };
77 Ok(Box::pin(stream::unfold(state, next_decoded::<R, C, P>)))
78 }
79}
80
81struct DecodeState<R, C, P>
82where
83 R: Reader<Payload>,
84{
85 inner: Pin<Box<R::Stream>>,
86 codec: C,
87 disposition: DecodeErrorDisposition,
88 _payload: PhantomData<P>,
89}
90
91async fn next_decoded<R, C, P>(
92 mut state: DecodeState<R, C, P>,
93) -> Option<(
94 Result<Message<R::Acker, R::Cursor, P>>,
95 DecodeState<R, C, P>,
96)>
97where
98 R: Reader<Payload>,
99 C: EventCodec<P>,
100{
101 let item = state.inner.as_mut().next().await?;
102 match item {
103 Ok(msg) => {
104 let (event, acker, cursor) = msg.into_parts();
105 match state.codec.decode(event) {
106 Ok(decoded) => Some((Ok(Message::new(decoded, acker, cursor)), state)),
107 Err(e) => {
108 let disposition_result = match state.disposition {
109 DecodeErrorDisposition::Surface => Ok(()),
110 DecodeErrorDisposition::AckInner => acker.ack().await,
111 DecodeErrorDisposition::NackInner => acker.nack().await,
112 };
113 match disposition_result {
114 Ok(()) => Some((Err(e), state)),
115 Err(ack_err) => Some((Err(ack_err), state)),
116 }
117 }
118 }
119 }
120 Err(e) => Some((Err(e), state)),
121 }
122}
123
124pub trait ReaderTypedExt: Reader<Payload> + Sized {
126 fn decode<P, C>(self, codec: C) -> DecodeReader<Self, PayloadEventCodec<C>, P>
127 where
128 C: PayloadCodec<P>,
129 {
130 DecodeReader::from_payload_codec(self, codec)
131 }
132
133 fn decode_event<P, C>(self, codec: C) -> DecodeReader<Self, C, P>
134 where
135 C: EventCodec<P>,
136 {
137 DecodeReader::new(self, codec)
138 }
139}
140
141impl<T: Reader<Payload> + Sized> ReaderTypedExt for T {}
142
143#[cfg(test)]
144mod tests {
145 use super::*;
146
147 use std::sync::{Arc, Mutex};
148
149 use futures::Stream;
150 use futures::StreamExt;
151 use futures::stream;
152
153 use crate::error::Error;
154 use crate::event::Event;
155 use crate::io::NoCursor;
156 use crate::io::Writer;
157 use crate::io::acker::NoopAcker;
158 use crate::io::writer::EncodeWriter;
159 use crate::io::writer::WriterTypedExt;
160 use crate::payload::Payload as WirePayload;
161
162 #[derive(Debug, Clone, PartialEq, Eq, serde::Serialize, serde::Deserialize)]
163 struct UserUpdated {
164 user_id: String,
165 }
166
167 #[derive(Debug, Clone, PartialEq, Eq, serde::Serialize, serde::Deserialize)]
168 struct UserDeleted {
169 user_id: String,
170 }
171
172 #[derive(Debug, Clone, PartialEq, Eq)]
173 enum DomainEvent {
174 UserUpdated(UserUpdated),
175 UserDeleted(UserDeleted),
176 }
177
178 #[derive(Debug, Clone, Copy)]
179 struct DomainEventCodec;
180
181 impl EventCodec<DomainEvent> for DomainEventCodec {
182 fn encode(&self, event: &Event<DomainEvent>) -> Result<Event<WirePayload>> {
183 match event.payload() {
184 DomainEvent::UserUpdated(payload) => event
185 .clone()
186 .map_payload(|_| payload.clone())
187 .encode_payload(&crate::JsonPayloadCodec),
188 DomainEvent::UserDeleted(payload) => event
189 .clone()
190 .map_payload(|_| payload.clone())
191 .encode_payload(&crate::JsonPayloadCodec),
192 }
193 }
194
195 fn decode(&self, event: Event<WirePayload>) -> Result<Event<DomainEvent>> {
196 match event.topic().as_str() {
197 "user.updated" => event
198 .decode_payload::<UserUpdated, _>(&crate::JsonPayloadCodec)
199 .map(|event| event.map_payload(DomainEvent::UserUpdated)),
200 "user.deleted" => event
201 .decode_payload::<UserDeleted, _>(&crate::JsonPayloadCodec)
202 .map(|event| event.map_payload(DomainEvent::UserDeleted)),
203 topic => Err(Error::InvalidPayload(format!(
204 "unsupported domain event topic: {topic}"
205 ))),
206 }
207 }
208 }
209
210 #[derive(Clone, Default)]
211 struct CapturingPayloadWriter {
212 events: Arc<Mutex<Vec<Event<WirePayload>>>>,
213 }
214
215 impl Writer<WirePayload> for CapturingPayloadWriter {
216 async fn write(&self, event: &Event<WirePayload>) -> Result<()> {
217 self.events.lock().unwrap().push(event.clone());
218 Ok(())
219 }
220 }
221
222 #[tokio::test]
223 async fn encode_writer_converts_typed_event_to_payload_event() {
224 let inner = CapturingPayloadWriter::default();
225 let captured = Arc::clone(&inner.events);
226 let writer =
227 EncodeWriter::<_, _, UserUpdated>::from_payload_codec(inner, crate::JsonPayloadCodec);
228
229 let event = Event::create(
230 "acme",
231 "/users",
232 "user.updated",
233 "thing-1",
234 UserUpdated {
235 user_id: "u-1".to_owned(),
236 },
237 )
238 .unwrap();
239
240 writer.write(&event).await.unwrap();
241
242 let events = captured.lock().unwrap();
243 assert_eq!(events.len(), 1);
244 assert_eq!(events[0].payload().content_type(), crate::ContentType::Json);
245 let decoded: UserUpdated = events[0].payload().to_json().unwrap();
246 assert_eq!(decoded.user_id, "u-1");
247 }
248
249 #[tokio::test]
250 async fn encode_writer_accepts_event_codec_for_domain_enum() {
251 let inner = CapturingPayloadWriter::default();
252 let captured = Arc::clone(&inner.events);
253 let writer = EncodeWriter::<_, _, DomainEvent>::new(inner, DomainEventCodec);
254
255 let event = Event::create(
256 "acme",
257 "/users",
258 "user.deleted",
259 "thing-1",
260 DomainEvent::UserDeleted(UserDeleted {
261 user_id: "u-1".to_owned(),
262 }),
263 )
264 .unwrap();
265
266 writer.write(&event).await.unwrap();
267
268 let events = captured.lock().unwrap();
269 assert_eq!(events.len(), 1);
270 assert_eq!(events[0].topic().as_str(), "user.deleted");
271 let decoded: UserDeleted = events[0].payload().to_json().unwrap();
272 assert_eq!(decoded.user_id, "u-1");
273 }
274
275 #[derive(Debug, Clone, Default)]
276 struct TestSub;
277
278 struct VecPayloadReader {
279 events: Vec<Event<WirePayload>>,
280 }
281
282 impl VecPayloadReader {
283 fn new(events: Vec<Event<WirePayload>>) -> Self {
284 Self { events }
285 }
286 }
287
288 impl Reader<WirePayload> for VecPayloadReader {
289 type Subscription = TestSub;
290 type Acker = NoopAcker;
291 type Cursor = NoCursor;
292 type Stream =
293 Pin<Box<dyn Stream<Item = Result<Message<NoopAcker, NoCursor, WirePayload>>> + Send>>;
294
295 async fn read(&self, _: Self::Subscription) -> Result<Self::Stream> {
296 let events = self
297 .events
298 .clone()
299 .into_iter()
300 .map(|event| Ok(Message::new(event, NoopAcker, NoCursor)))
301 .collect::<Vec<_>>();
302 Ok(Box::pin(stream::iter(events)))
303 }
304 }
305
306 #[tokio::test]
307 async fn decode_reader_converts_payload_event_to_typed_event() {
308 let raw = VecPayloadReader::new(vec![
309 Event::create(
310 "acme",
311 "/users",
312 "user.updated",
313 "thing-1",
314 WirePayload::from_json(&UserUpdated {
315 user_id: "u-1".to_owned(),
316 })
317 .unwrap(),
318 )
319 .unwrap(),
320 ]);
321
322 let reader =
323 DecodeReader::<_, _, UserUpdated>::from_payload_codec(raw, crate::JsonPayloadCodec);
324 let mut stream = reader.read(TestSub).await.unwrap();
325 let msg = stream.next().await.unwrap().unwrap();
326
327 assert_eq!(msg.event().payload().user_id, "u-1");
328 }
329
330 #[derive(Debug, Clone, Default)]
331 struct CountingAcker {
332 ack_count: Arc<std::sync::atomic::AtomicUsize>,
333 nack_count: Arc<std::sync::atomic::AtomicUsize>,
334 }
335
336 impl Acker for CountingAcker {
337 async fn ack(&self) -> Result<()> {
338 self.ack_count
339 .fetch_add(1, std::sync::atomic::Ordering::SeqCst);
340 Ok(())
341 }
342 async fn nack(&self) -> Result<()> {
343 self.nack_count
344 .fetch_add(1, std::sync::atomic::Ordering::SeqCst);
345 Ok(())
346 }
347 }
348
349 #[tokio::test]
350 async fn decode_reader_acks_inner_on_decode_failure_by_default() {
351 let acker = CountingAcker::default();
352 let raw = {
353 struct TypedVecReader {
354 events: Vec<Message<CountingAcker, NoCursor, WirePayload>>,
355 }
356 impl Reader<WirePayload> for TypedVecReader {
357 type Subscription = TestSub;
358 type Acker = CountingAcker;
359 type Cursor = NoCursor;
360 type Stream = Pin<
361 Box<
362 dyn Stream<Item = Result<Message<CountingAcker, NoCursor, WirePayload>>>
363 + Send,
364 >,
365 >;
366
367 async fn read(&self, _: Self::Subscription) -> Result<Self::Stream> {
368 let items: Vec<Result<Message<CountingAcker, NoCursor, WirePayload>>> =
369 self.events.clone().into_iter().map(Ok).collect();
370 Ok(Box::pin(stream::iter(items)))
371 }
372 }
373 TypedVecReader {
374 events: vec![Message::new(
375 Event::create(
376 "acme",
377 "/users",
378 "user.updated",
379 "thing-1",
380 WirePayload::from_string("bad"),
381 )
382 .unwrap(),
383 acker.clone(),
384 NoCursor,
385 )],
386 }
387 };
388
389 let reader =
390 DecodeReader::<_, _, UserUpdated>::from_payload_codec(raw, crate::JsonPayloadCodec);
391 let mut stream = reader.read(TestSub).await.unwrap();
392 let err = stream.next().await.unwrap().unwrap_err();
393
394 assert!(matches!(
395 err,
396 Error::InvalidPayload(_) | Error::Serialization(_)
397 ));
398 assert_eq!(acker.ack_count.load(std::sync::atomic::Ordering::SeqCst), 1);
399 assert_eq!(
400 acker.nack_count.load(std::sync::atomic::Ordering::SeqCst),
401 0
402 );
403 }
404
405 #[tokio::test]
406 async fn decode_reader_can_nack_inner_on_decode_failure() {
407 let acker = CountingAcker::default();
408 let raw = {
409 struct TypedVecReader {
410 events: Vec<Message<CountingAcker, NoCursor, WirePayload>>,
411 }
412 impl Reader<WirePayload> for TypedVecReader {
413 type Subscription = TestSub;
414 type Acker = CountingAcker;
415 type Cursor = NoCursor;
416 type Stream = Pin<
417 Box<
418 dyn Stream<Item = Result<Message<CountingAcker, NoCursor, WirePayload>>>
419 + Send,
420 >,
421 >;
422
423 async fn read(&self, _: Self::Subscription) -> Result<Self::Stream> {
424 let items: Vec<Result<Message<CountingAcker, NoCursor, WirePayload>>> =
425 self.events.clone().into_iter().map(Ok).collect();
426 Ok(Box::pin(stream::iter(items)))
427 }
428 }
429 TypedVecReader {
430 events: vec![Message::new(
431 Event::create(
432 "acme",
433 "/users",
434 "user.updated",
435 "thing-1",
436 WirePayload::from_string("bad"),
437 )
438 .unwrap(),
439 acker.clone(),
440 NoCursor,
441 )],
442 }
443 };
444
445 let reader =
446 DecodeReader::<_, _, UserUpdated>::from_payload_codec(raw, crate::JsonPayloadCodec)
447 .with_disposition(DecodeErrorDisposition::NackInner);
448 let mut stream = reader.read(TestSub).await.unwrap();
449 let err = stream.next().await.unwrap().unwrap_err();
450
451 assert!(matches!(
452 err,
453 Error::InvalidPayload(_) | Error::Serialization(_)
454 ));
455 assert_eq!(acker.ack_count.load(std::sync::atomic::Ordering::SeqCst), 0);
456 assert_eq!(
457 acker.nack_count.load(std::sync::atomic::Ordering::SeqCst),
458 1
459 );
460 }
461
462 #[tokio::test]
463 async fn event_codec_reader_decodes_domain_enum_pipeline() {
464 let raw = VecPayloadReader::new(vec![
465 Event::create(
466 "acme",
467 "/users",
468 "user.deleted",
469 "thing-1",
470 WirePayload::from_json(&UserDeleted {
471 user_id: "u-1".to_owned(),
472 })
473 .unwrap(),
474 )
475 .unwrap(),
476 ]);
477
478 let reader = DecodeReader::<_, _, DomainEvent>::new(raw, DomainEventCodec);
479 let mut stream = reader.read(TestSub).await.unwrap();
480 let msg = stream.next().await.unwrap().unwrap();
481
482 assert_eq!(
483 msg.event().payload(),
484 &DomainEvent::UserDeleted(UserDeleted {
485 user_id: "u-1".to_owned()
486 })
487 );
488 }
489
490 #[tokio::test]
491 async fn writer_typed_ext_encodes_with_payload_codec() {
492 let inner = CapturingPayloadWriter::default();
493 let captured = Arc::clone(&inner.events);
494 let writer = inner.encode::<UserUpdated, _>(crate::JsonPayloadCodec);
495
496 writer
497 .write(
498 &Event::create(
499 "acme",
500 "/users",
501 "user.updated",
502 "thing-1",
503 UserUpdated {
504 user_id: "u-1".to_owned(),
505 },
506 )
507 .unwrap(),
508 )
509 .await
510 .unwrap();
511
512 let raw = captured.lock().unwrap();
513 let decoded: UserUpdated = raw[0].payload().to_json().unwrap();
514 assert_eq!(decoded.user_id, "u-1");
515 }
516
517 #[tokio::test]
518 async fn reader_typed_ext_decodes_with_payload_codec() {
519 let raw = VecPayloadReader::new(vec![
520 Event::create(
521 "acme",
522 "/users",
523 "user.updated",
524 "thing-1",
525 WirePayload::from_json(&UserUpdated {
526 user_id: "u-1".to_owned(),
527 })
528 .unwrap(),
529 )
530 .unwrap(),
531 ]);
532
533 let reader = raw.decode::<UserUpdated, _>(crate::JsonPayloadCodec);
534 let mut stream = reader.read(TestSub).await.unwrap();
535 let msg = stream.next().await.unwrap().unwrap();
536
537 assert_eq!(msg.event().payload().user_id, "u-1");
538 }
539
540 #[tokio::test]
541 async fn typed_writer_encodes_to_payload_writer() {
542 let inner = CapturingPayloadWriter::default();
543 let captured = Arc::clone(&inner.events);
544 let writer = inner.encode::<UserUpdated, _>(crate::JsonPayloadCodec);
545
546 writer
547 .write(
548 &Event::create(
549 "acme",
550 "/users",
551 "user.updated",
552 "thing-1",
553 UserUpdated {
554 user_id: "u-1".to_owned(),
555 },
556 )
557 .unwrap(),
558 )
559 .await
560 .unwrap();
561
562 let raw = captured.lock().unwrap();
563 let decoded: UserUpdated = raw[0].payload().to_json().unwrap();
564 assert_eq!(decoded.user_id, "u-1");
565 }
566}