1use std::{
35 any::{Any, TypeId},
36 collections::HashMap,
37 fmt::Debug,
38 marker::PhantomData,
39 sync::Arc,
40};
41
42use nautilus_core::UUID4;
43
44use crate::{
45 capture::encoder::{Encode, EncodeError, EncodedPayload, TypedEncoder},
46 entry::PayloadType,
47 headers::Headers,
48};
49
50pub trait HeadersExtractor: Send + Sync {
59 fn extract(&self, message: &dyn Any) -> Headers;
63}
64
65pub struct TypedHeadersExtractor<T: 'static, F> {
67 func: F,
68 _phantom: PhantomData<fn(&T)>,
69}
70
71impl<T: 'static, F> TypedHeadersExtractor<T, F>
72where
73 F: Fn(&T) -> Headers + Send + Sync,
74{
75 #[must_use]
77 pub const fn new(func: F) -> Self {
78 Self {
79 func,
80 _phantom: PhantomData,
81 }
82 }
83}
84
85impl<T: 'static, F> Debug for TypedHeadersExtractor<T, F> {
86 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
87 f.debug_struct(stringify!(TypedHeadersExtractor))
88 .field("type", &std::any::type_name::<T>())
89 .finish_non_exhaustive()
90 }
91}
92
93impl<T: 'static, F> HeadersExtractor for TypedHeadersExtractor<T, F>
94where
95 F: Fn(&T) -> Headers + Send + Sync,
96{
97 fn extract(&self, message: &dyn Any) -> Headers {
98 message
99 .downcast_ref::<T>()
100 .map(&self.func)
101 .unwrap_or_default()
102 }
103}
104
105#[derive(Debug, Default)]
108struct EmptyHeadersExtractor;
109
110impl HeadersExtractor for EmptyHeadersExtractor {
111 fn extract(&self, _: &dyn Any) -> Headers {
112 Headers::empty()
113 }
114}
115
116pub trait IdentityExtractor: Send + Sync {
122 fn extract(&self, message: &dyn Any) -> Option<UUID4>;
125}
126
127pub struct TypedIdentityExtractor<T: 'static, F> {
129 func: F,
130 _phantom: PhantomData<fn(&T)>,
131}
132
133impl<T: 'static, F> TypedIdentityExtractor<T, F>
134where
135 F: Fn(&T) -> Option<UUID4> + Send + Sync,
136{
137 #[must_use]
139 pub const fn new(func: F) -> Self {
140 Self {
141 func,
142 _phantom: PhantomData,
143 }
144 }
145}
146
147impl<T: 'static, F> Debug for TypedIdentityExtractor<T, F> {
148 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
149 f.debug_struct(stringify!(TypedIdentityExtractor))
150 .field("type", &std::any::type_name::<T>())
151 .finish_non_exhaustive()
152 }
153}
154
155impl<T: 'static, F> IdentityExtractor for TypedIdentityExtractor<T, F>
156where
157 F: Fn(&T) -> Option<UUID4> + Send + Sync,
158{
159 fn extract(&self, message: &dyn Any) -> Option<UUID4> {
160 message.downcast_ref::<T>().and_then(&self.func)
161 }
162}
163
164#[derive(Clone)]
167struct Registered {
168 payload_type: PayloadType,
169 encoder: Arc<dyn Encode>,
170 headers: Arc<dyn HeadersExtractor>,
171 identity: Option<Arc<dyn IdentityExtractor>>,
172}
173
174impl Debug for Registered {
175 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
176 f.debug_struct(stringify!(Registered))
177 .field("payload_type", &self.payload_type.as_str())
178 .finish_non_exhaustive()
179 }
180}
181
182#[derive(Clone, Debug, Default)]
189pub struct EncoderRegistry {
190 by_type: HashMap<TypeId, Registered>,
191}
192
193impl EncoderRegistry {
194 #[must_use]
196 pub fn new() -> Self {
197 Self::default()
198 }
199
200 pub fn register<T, F>(&mut self, payload_type: PayloadType, func: F)
208 where
209 T: 'static,
210 F: Fn(&T) -> Result<EncodedPayload, EncodeError> + Send + Sync + 'static,
211 {
212 let encoder: Arc<dyn Encode> = Arc::new(TypedEncoder::<T, F>::new(func));
213 let headers = self
214 .preserved_headers::<T>()
215 .unwrap_or_else(|| Arc::new(EmptyHeadersExtractor) as Arc<dyn HeadersExtractor>);
216 let identity = self.preserved_identity::<T>();
217 self.by_type.insert(
218 TypeId::of::<T>(),
219 Registered {
220 payload_type,
221 encoder,
222 headers,
223 identity,
224 },
225 );
226 }
227
228 pub fn register_with_headers<T, F, H>(
231 &mut self,
232 payload_type: PayloadType,
233 func: F,
234 headers_fn: H,
235 ) where
236 T: 'static,
237 F: Fn(&T) -> Result<EncodedPayload, EncodeError> + Send + Sync + 'static,
238 H: Fn(&T) -> Headers + Send + Sync + 'static,
239 {
240 let encoder: Arc<dyn Encode> = Arc::new(TypedEncoder::<T, F>::new(func));
241 let headers: Arc<dyn HeadersExtractor> =
242 Arc::new(TypedHeadersExtractor::<T, H>::new(headers_fn));
243 let identity = self.preserved_identity::<T>();
244 self.by_type.insert(
245 TypeId::of::<T>(),
246 Registered {
247 payload_type,
248 encoder,
249 headers,
250 identity,
251 },
252 );
253 }
254
255 pub fn register_encoder<T: 'static>(
261 &mut self,
262 payload_type: PayloadType,
263 encoder: Arc<dyn Encode>,
264 ) {
265 let headers = self
266 .preserved_headers::<T>()
267 .unwrap_or_else(|| Arc::new(EmptyHeadersExtractor) as Arc<dyn HeadersExtractor>);
268 let identity = self.preserved_identity::<T>();
269 self.by_type.insert(
270 TypeId::of::<T>(),
271 Registered {
272 payload_type,
273 encoder,
274 headers,
275 identity,
276 },
277 );
278 }
279
280 pub fn register_headers<T, H>(&mut self, headers_fn: H)
287 where
288 T: 'static,
289 H: Fn(&T) -> Headers + Send + Sync + 'static,
290 {
291 if let Some(reg) = self.by_type.get_mut(&TypeId::of::<T>()) {
292 reg.headers = Arc::new(TypedHeadersExtractor::<T, H>::new(headers_fn));
293 }
294 }
295
296 pub fn register_identity<T, F>(&mut self, identity_fn: F)
302 where
303 T: 'static,
304 F: Fn(&T) -> Option<UUID4> + Send + Sync + 'static,
305 {
306 if let Some(reg) = self.by_type.get_mut(&TypeId::of::<T>()) {
307 reg.identity = Some(Arc::new(TypedIdentityExtractor::<T, F>::new(identity_fn)));
308 }
309 }
310
311 fn preserved_headers<T: 'static>(&self) -> Option<Arc<dyn HeadersExtractor>> {
312 self.by_type
313 .get(&TypeId::of::<T>())
314 .map(|reg| Arc::clone(®.headers))
315 }
316
317 fn preserved_identity<T: 'static>(&self) -> Option<Arc<dyn IdentityExtractor>> {
318 self.by_type
319 .get(&TypeId::of::<T>())
320 .and_then(|reg| reg.identity.clone())
321 }
322
323 #[must_use]
325 pub fn len(&self) -> usize {
326 self.by_type.len()
327 }
328
329 #[must_use]
331 pub fn is_empty(&self) -> bool {
332 self.by_type.is_empty()
333 }
334
335 #[must_use]
337 pub fn contains<T: 'static>(&self) -> bool {
338 self.by_type.contains_key(&TypeId::of::<T>())
339 }
340
341 pub fn encode<T: 'static>(
353 &self,
354 message: &T,
355 ) -> Result<Option<(PayloadType, EncodedPayload)>, EncodeError> {
356 let Some(reg) = self.by_type.get(&TypeId::of::<T>()) else {
357 return Ok(None);
358 };
359
360 let encoded = reg.encoder.encode(message as &dyn Any)?;
361 let payload_type = encoded.payload_type.unwrap_or(reg.payload_type);
362 Ok(Some((payload_type, encoded)))
363 }
364
365 pub fn encode_any(
376 &self,
377 message: &dyn Any,
378 ) -> Result<Option<(PayloadType, EncodedPayload)>, EncodeError> {
379 let Some(reg) = self.by_type.get(&message.type_id()) else {
380 return Ok(None);
381 };
382
383 let encoded = reg.encoder.encode(message)?;
384 let payload_type = encoded.payload_type.unwrap_or(reg.payload_type);
385 Ok(Some((payload_type, encoded)))
386 }
387
388 #[must_use]
392 pub fn headers_for_any(&self, message: &dyn Any) -> Option<Headers> {
393 self.by_type
394 .get(&message.type_id())
395 .map(|reg| reg.headers.extract(message))
396 }
397
398 #[must_use]
402 pub fn identity_for_any(&self, message: &dyn Any) -> Option<UUID4> {
403 self.by_type
404 .get(&message.type_id())
405 .and_then(|reg| reg.identity.as_ref())
406 .and_then(|identity| identity.extract(message))
407 }
408}
409
410#[cfg(test)]
411mod tests {
412 use std::sync::Arc;
413
414 use bytes::Bytes;
415 use rstest::rstest;
416 use ustr::Ustr;
417
418 use super::*;
419
420 #[derive(Debug)]
421 struct Sample(u8);
422
423 #[derive(Debug)]
424 struct Other;
425
426 #[derive(Debug)]
427 struct StatefulEncoder {
428 prefix: u8,
429 }
430
431 impl Encode for StatefulEncoder {
432 fn encode(&self, message: &dyn Any) -> Result<EncodedPayload, EncodeError> {
433 let sample = message
434 .downcast_ref::<Sample>()
435 .ok_or(EncodeError::TypeMismatch {
436 expected: std::any::type_name::<Sample>(),
437 })?;
438 Ok(EncodedPayload::without_indices(Bytes::copy_from_slice(&[
439 self.prefix,
440 sample.0,
441 ])))
442 }
443 }
444
445 #[rstest]
446 fn unknown_type_returns_none() {
447 let registry = EncoderRegistry::new();
448
449 assert!(registry.encode(&Sample(1)).expect("encode").is_none());
450 assert!(!registry.contains::<Sample>());
451 }
452
453 #[rstest]
454 fn registered_type_returns_payload_type_and_payload() {
455 let mut registry = EncoderRegistry::new();
456 registry.register::<Sample, _>(Ustr::from("Sample"), |s| {
457 Ok(EncodedPayload::without_indices(Bytes::copy_from_slice(&[
458 s.0,
459 ])))
460 });
461
462 let (tag, encoded) = registry.encode(&Sample(9)).expect("encode").expect("hit");
463
464 assert_eq!(tag.as_str(), "Sample");
465 assert_eq!(encoded.payload.as_ref(), &[9]);
466 assert!(registry.contains::<Sample>());
467 assert_eq!(registry.len(), 1);
468 }
469
470 #[rstest]
471 fn re_registering_replaces_prior_encoder() {
472 let mut registry = EncoderRegistry::new();
473 registry.register::<Sample, _>(Ustr::from("Old"), |s| {
474 Ok(EncodedPayload::without_indices(Bytes::copy_from_slice(&[
475 s.0,
476 ])))
477 });
478 registry.register::<Sample, _>(Ustr::from("New"), |s| {
479 Ok(EncodedPayload::without_indices(Bytes::copy_from_slice(&[
480 s.0, s.0,
481 ])))
482 });
483
484 let (tag, encoded) = registry.encode(&Sample(3)).expect("encode").expect("hit");
485
486 assert_eq!(tag.as_str(), "New");
487 assert_eq!(encoded.payload.as_ref(), &[3, 3]);
488 assert_eq!(registry.len(), 1);
489 }
490
491 #[rstest]
492 fn register_encoder_replaces_encoder_and_preserves_extractors() {
493 let mut registry = EncoderRegistry::new();
494 registry.register::<Sample, _>(Ustr::from("Old"), |s| {
495 Ok(EncodedPayload::without_indices(Bytes::copy_from_slice(&[
496 s.0,
497 ])))
498 });
499 let correlation_id = UUID4::new();
500 let causation_id = UUID4::new();
501 registry.register_headers::<Sample, _>(move |_| Headers {
502 correlation_id: Some(correlation_id),
503 causation_id: Some(causation_id),
504 });
505 let identity = UUID4::new();
506 registry.register_identity::<Sample, _>(move |_| Some(identity));
507
508 registry.register_encoder::<Sample>(
509 Ustr::from("Stateful"),
510 Arc::new(StatefulEncoder { prefix: 0xA5 }),
511 );
512 let sample = Sample(7);
513 let (payload_type, encoded) = registry.encode(&sample).expect("encode").expect("hit");
514 let headers = registry
515 .headers_for_any(&sample as &dyn Any)
516 .expect("headers");
517
518 assert_eq!(payload_type.as_str(), "Stateful");
519 assert_eq!(
520 encoded,
521 EncodedPayload::without_indices(Bytes::from_static(&[0xA5, 7])),
522 );
523 assert_eq!(
524 headers,
525 Headers {
526 correlation_id: Some(correlation_id),
527 causation_id: Some(causation_id),
528 },
529 );
530 assert_eq!(
531 registry.identity_for_any(&sample as &dyn Any),
532 Some(identity)
533 );
534 assert_eq!(registry.len(), 1);
535 }
536
537 #[rstest]
538 fn registry_is_empty_by_default() {
539 let registry = EncoderRegistry::new();
540
541 assert!(registry.is_empty());
542 assert_eq!(registry.len(), 0);
543 assert!(!registry.contains::<Other>());
544 }
545
546 #[rstest]
547 fn encode_any_dispatches_by_concrete_type_id() {
548 let mut registry = EncoderRegistry::new();
552 registry.register::<Sample, _>(Ustr::from("Sample"), |s| {
553 Ok(EncodedPayload::without_indices(Bytes::copy_from_slice(&[
554 s.0,
555 ])))
556 });
557
558 let sample = Sample(5);
559 let (tag, encoded) = registry
560 .encode_any(&sample as &dyn Any)
561 .expect("encode_any")
562 .expect("hit");
563
564 assert_eq!(tag.as_str(), "Sample");
565 assert_eq!(encoded.payload.as_ref(), &[5]);
566 }
567
568 #[rstest]
569 fn encode_any_returns_none_for_unregistered_type() {
570 let registry = EncoderRegistry::new();
574
575 let unregistered = Other;
576 let outcome = registry
577 .encode_any(&unregistered as &dyn Any)
578 .expect("encode_any");
579
580 assert!(outcome.is_none());
581 }
582
583 #[rstest]
584 fn encoder_payload_type_override_overrides_registered_tag() {
585 let mut registry = EncoderRegistry::new();
591 registry.register::<Sample, _>(Ustr::from("Wrapper"), |s| {
592 Ok(EncodedPayload::with_payload_type(
593 Ustr::from("Inner"),
594 Bytes::copy_from_slice(&[s.0]),
595 Vec::new(),
596 ))
597 });
598
599 let (tag, _) = registry.encode(&Sample(1)).expect("encode").expect("hit");
600 assert_eq!(tag.as_str(), "Inner");
601
602 let (any_tag, _) = registry
603 .encode_any(&Sample(1) as &dyn Any)
604 .expect("encode_any")
605 .expect("hit");
606 assert_eq!(any_tag.as_str(), "Inner");
607 }
608
609 #[rstest]
610 fn registered_type_without_headers_extractor_returns_empty_headers() {
611 let mut registry = EncoderRegistry::new();
612 registry.register::<Sample, _>(Ustr::from("Sample"), |s| {
613 Ok(EncodedPayload::without_indices(Bytes::copy_from_slice(&[
614 s.0,
615 ])))
616 });
617
618 let headers = registry
619 .headers_for_any(&Sample(1) as &dyn Any)
620 .expect("hit");
621 assert_eq!(headers, Headers::empty());
622 }
623
624 #[rstest]
625 fn headers_for_any_returns_none_for_unregistered_type() {
626 let registry = EncoderRegistry::new();
627 let outcome = registry.headers_for_any(&Other as &dyn Any);
628
629 assert!(outcome.is_none());
630 }
631
632 #[rstest]
633 fn register_with_headers_uses_extractor() {
634 let mut registry = EncoderRegistry::new();
638 let causation = nautilus_core::UUID4::new();
639 let causation_captured = causation;
640 registry.register_with_headers::<Sample, _, _>(
641 Ustr::from("Sample"),
642 |s| {
643 Ok(EncodedPayload::without_indices(Bytes::copy_from_slice(&[
644 s.0,
645 ])))
646 },
647 move |_| Headers {
648 correlation_id: None,
649 causation_id: Some(causation_captured),
650 },
651 );
652
653 let headers = registry
654 .headers_for_any(&Sample(1) as &dyn Any)
655 .expect("hit");
656 assert_eq!(headers.causation_id, Some(causation));
657 }
658
659 #[rstest]
660 fn register_headers_overrides_default_extractor_post_register() {
661 let mut registry = EncoderRegistry::new();
665 registry.register::<Sample, _>(Ustr::from("Sample"), |s| {
666 Ok(EncodedPayload::without_indices(Bytes::copy_from_slice(&[
667 s.0,
668 ])))
669 });
670 let correlation = nautilus_core::UUID4::new();
671 let correlation_captured = correlation;
672 registry.register_headers::<Sample, _>(move |_| Headers {
673 correlation_id: Some(correlation_captured),
674 causation_id: None,
675 });
676
677 let headers = registry
678 .headers_for_any(&Sample(1) as &dyn Any)
679 .expect("hit");
680 assert_eq!(headers.correlation_id, Some(correlation));
681 }
682
683 #[rstest]
684 fn identity_for_any_returns_none_without_extractor() {
685 let mut registry = EncoderRegistry::new();
686 registry.register::<Sample, _>(Ustr::from("Sample"), |s| {
687 Ok(EncodedPayload::without_indices(Bytes::copy_from_slice(&[
688 s.0,
689 ])))
690 });
691
692 assert!(registry.identity_for_any(&Sample(1) as &dyn Any).is_none());
693 assert!(registry.identity_for_any(&Other as &dyn Any).is_none());
694 }
695
696 #[rstest]
697 fn register_identity_extracts_and_survives_re_register() {
698 let mut registry = EncoderRegistry::new();
701 registry.register::<Sample, _>(Ustr::from("Old"), |s| {
702 Ok(EncodedPayload::without_indices(Bytes::copy_from_slice(&[
703 s.0,
704 ])))
705 });
706 let identity = nautilus_core::UUID4::new();
707 registry.register_identity::<Sample, _>(move |_| Some(identity));
708
709 assert_eq!(
710 registry.identity_for_any(&Sample(1) as &dyn Any),
711 Some(identity),
712 );
713
714 registry.register::<Sample, _>(Ustr::from("New"), |s| {
715 Ok(EncodedPayload::without_indices(Bytes::copy_from_slice(&[
716 s.0, s.0,
717 ])))
718 });
719
720 assert_eq!(
721 registry.identity_for_any(&Sample(1) as &dyn Any),
722 Some(identity),
723 );
724 }
725
726 #[rstest]
727 fn register_identity_for_unregistered_type_is_silent_noop() {
728 let mut registry = EncoderRegistry::new();
729 registry.register_identity::<Sample, _>(|_| Some(nautilus_core::UUID4::new()));
730
731 assert!(!registry.contains::<Sample>());
732 assert!(registry.identity_for_any(&Sample(1) as &dyn Any).is_none());
733 }
734
735 #[rstest]
736 fn register_headers_for_unregistered_type_is_silent_noop() {
737 let mut registry = EncoderRegistry::new();
738 registry.register_headers::<Sample, _>(|_| Headers::empty());
739
740 assert!(!registry.contains::<Sample>());
741 assert!(registry.headers_for_any(&Sample(1) as &dyn Any).is_none());
742 }
743
744 #[rstest]
745 fn re_registering_preserves_existing_headers_extractor() {
746 let mut registry = EncoderRegistry::new();
750 registry.register::<Sample, _>(Ustr::from("Old"), |s| {
751 Ok(EncodedPayload::without_indices(Bytes::copy_from_slice(&[
752 s.0,
753 ])))
754 });
755 let causation = nautilus_core::UUID4::new();
756 let causation_captured = causation;
757 registry.register_headers::<Sample, _>(move |_| Headers {
758 correlation_id: None,
759 causation_id: Some(causation_captured),
760 });
761 registry.register::<Sample, _>(Ustr::from("New"), |s| {
762 Ok(EncodedPayload::without_indices(Bytes::copy_from_slice(&[
763 s.0, s.0,
764 ])))
765 });
766
767 let (tag, _) = registry.encode(&Sample(3)).expect("encode").expect("hit");
768 assert_eq!(tag.as_str(), "New");
769 let headers = registry
770 .headers_for_any(&Sample(3) as &dyn Any)
771 .expect("hit");
772 assert_eq!(headers.causation_id, Some(causation));
773 }
774}