use std::any::TypeId; use std::collections::HashSet; use proptest::prelude::*; use swactor::actor::ActorAddress; use swactor::Error; use swactor_transport::{ Codec, CodecRegistrationError, CodecRegistry, JsonCodec, NetworkMessage, WireEnvelope, }; #[derive(Clone, Debug, PartialEq, Eq)] struct Number(u64); impl NetworkMessage for Number { fn type_tag() -> &'static str { "contract::Number" } } #[derive(Clone, Debug, PartialEq, Eq)] struct Text(String); impl NetworkMessage for Text { fn type_tag() -> &'static str { "contract::Text" } } #[derive(Clone, Debug, PartialEq, Eq)] struct Blob(Vec); impl NetworkMessage for Blob { fn type_tag() -> &'static str { "contract::Blob" } } #[derive(Clone, Debug, PartialEq, Eq, serde::Serialize, serde::Deserialize)] struct Structured { id: u64, labels: Vec, } impl NetworkMessage for Structured { fn type_tag() -> &'static str { "contract::Structured" } } #[derive(Clone, Debug, PartialEq, Eq)] struct Fallible(u8); impl NetworkMessage for Fallible { fn type_tag() -> &'static str { "contract::Fallible" } } #[derive(Clone, Copy)] struct NumberCodec; impl Codec for NumberCodec { fn encode(&self, msg: &Number) -> Result, Error> { Ok(msg.0.to_be_bytes().to_vec()) } fn decode(&self, bytes: &[u8]) -> Result { let bytes: [u8; 8] = bytes .try_into() .map_err(|_| Error::from("number payload must contain eight bytes"))?; Ok(Number(u64::from_be_bytes(bytes))) } } #[derive(Clone, Copy)] struct TextCodec; impl Codec for TextCodec { fn encode(&self, msg: &Text) -> Result, Error> { Ok(msg.0.as_bytes().to_vec()) } fn decode(&self, bytes: &[u8]) -> Result { let text = std::str::from_utf8(bytes) .map_err(|error| Error::from(format!("invalid text payload: {error}")))?; Ok(Text(text.to_owned())) } } #[derive(Clone, Copy)] struct BlobCodec; impl Codec for BlobCodec { fn encode(&self, msg: &Blob) -> Result, Error> { Ok(msg.0.clone()) } fn decode(&self, bytes: &[u8]) -> Result { Ok(Blob(bytes.to_vec())) } } #[derive(Clone, Copy)] struct FallibleCodec; impl Codec for FallibleCodec { fn encode(&self, msg: &Fallible) -> Result, Error> { if msg.0 == u8::MAX { return Err(Error::from("refused value")); } Ok(vec![msg.0]) } fn decode(&self, bytes: &[u8]) -> Result { match bytes { [value] if *value != u8::MAX => Ok(Fallible(*value)), _ => Err(Error::from("malformed fallible payload")), } } } #[derive(Clone, Debug, PartialEq, Eq)] enum Sample { Number(Number), Text(Text), Blob(Blob), Structured(Structured), Fallible(Fallible), } impl Sample { fn register(&self, registry: &mut CodecRegistry) -> Result<(), CodecRegistrationError> { match self { Self::Number(_) => registry.register::(NumberCodec), Self::Text(_) => registry.register::(TextCodec), Self::Blob(_) => registry.register::(BlobCodec), Self::Structured(_) => registry.register::(JsonCodec::default()), Self::Fallible(_) => registry.register::(FallibleCodec), } } fn type_id(&self) -> TypeId { match self { Self::Number(_) => TypeId::of::(), Self::Text(_) => TypeId::of::(), Self::Blob(_) => TypeId::of::(), Self::Structured(_) => TypeId::of::(), Self::Fallible(_) => TypeId::of::(), } } fn tag(&self) -> &'static str { match self { Self::Number(_) => Number::type_tag(), Self::Text(_) => Text::type_tag(), Self::Blob(_) => Blob::type_tag(), Self::Structured(_) => Structured::type_tag(), Self::Fallible(_) => Fallible::type_tag(), } } fn boxed(&self) -> Box { match self { Self::Number(value) => Box::new(value.clone()), Self::Text(value) => Box::new(value.clone()), Self::Blob(value) => Box::new(value.clone()), Self::Structured(value) => Box::new(value.clone()), Self::Fallible(value) => Box::new(value.clone()), } } fn assert_decoded(&self, decoded: Box) { match self { Self::Number(expected) => assert_eq!(*decoded.downcast::().unwrap(), *expected), Self::Text(expected) => assert_eq!(*decoded.downcast::().unwrap(), *expected), Self::Blob(expected) => assert_eq!(*decoded.downcast::().unwrap(), *expected), Self::Structured(expected) => { assert_eq!(*decoded.downcast::().unwrap(), *expected) } Self::Fallible(expected) => { assert_eq!(*decoded.downcast::().unwrap(), *expected) } } } } fn address(seed: u8) -> ActorAddress { ActorAddress([seed; 32]) } fn assert_sample(registry: &CodecRegistry, sample: &Sample, destination: ActorAddress) -> Vec { let (tag, payload) = registry .encode(sample.type_id(), sample.boxed()) .expect("registered sample encodes"); assert_eq!(tag, sample.tag()); sample.assert_decoded( registry .decode(&tag, &payload) .expect("registered sample decodes"), ); let (received_destination, decoded) = registry .receive(WireEnvelope { dest: destination, type_tag: tag, payload: payload.clone(), }) .expect("valid envelope receives"); assert_eq!(received_destination, destination); sample.assert_decoded(decoded); payload } fn fingerprint(registry: &CodecRegistry, samples: &[Sample]) -> Vec> { let type_ids: HashSet = samples.iter().map(Sample::type_id).collect(); let tags: HashSet<&str> = samples.iter().map(Sample::tag).collect(); assert_eq!( type_ids.len(), samples.len(), "registered TypeIds are unique" ); assert_eq!(tags.len(), samples.len(), "registered wire tags are unique"); samples .iter() .enumerate() .map(|(index, sample)| assert_sample(registry, sample, address(index as u8))) .collect() } #[derive(Clone, Debug)] enum LegalOperation { Encode(usize), Decode(usize), Receive(usize, u8), } fn legal_operation() -> impl Strategy { prop_oneof![ any::().prop_map(LegalOperation::Encode), any::().prop_map(LegalOperation::Decode), (any::(), any::()).prop_map(|(slot, dest)| LegalOperation::Receive(slot, dest)), ] } proptest! { #![proptest_config(ProptestConfig { cases: 32, max_shrink_iters: 10_000, .. ProptestConfig::default() })] #[test] fn long_legal_action_sequences_preserve_every_registration( number in any::(), text in any::(), blob in prop::collection::vec(any::(), 0..512), structured_id in any::(), labels in prop::collection::vec(any::(), 0..16), fallible in 0u8..u8::MAX, operations in prop::collection::vec(legal_operation(), 128..1025), ) { let samples = vec![ Sample::Number(Number(number)), Sample::Text(Text(text)), Sample::Blob(Blob(blob)), Sample::Structured(Structured { id: structured_id, labels }), Sample::Fallible(Fallible(fallible)), ]; let mut registry = CodecRegistry::new(); let mut registered = Vec::new(); for sample in &samples { sample.register(&mut registry).expect("fresh type and tag register"); registered.push(sample.clone()); fingerprint(®istry, ®istered); } for operation in operations { let slot = match operation { LegalOperation::Encode(slot) | LegalOperation::Decode(slot) | LegalOperation::Receive(slot, _) => slot % samples.len(), }; let sample = &samples[slot]; match operation { LegalOperation::Encode(_) => { let (tag, _) = registry.encode(sample.type_id(), sample.boxed()).unwrap(); prop_assert_eq!(tag, sample.tag()); } LegalOperation::Decode(_) => { let (_, payload) = registry.encode(sample.type_id(), sample.boxed()).unwrap(); sample.assert_decoded(registry.decode(sample.tag(), &payload).unwrap()); } LegalOperation::Receive(_, dest) => { assert_sample(®istry, sample, address(dest)); } } fingerprint(®istry, ®istered); } } } #[test] fn every_registration_duplicate_is_rejected_without_replacement() { let sample = Sample::Number(Number(7)); let mut registry = CodecRegistry::new(); sample.register(&mut registry).unwrap(); let before = fingerprint(®istry, std::slice::from_ref(&sample)); assert!(matches!( registry.register::(NumberCodec), Err(CodecRegistrationError::EncoderAlreadyRegistered { .. }) )); assert_eq!( fingerprint(®istry, std::slice::from_ref(&sample)), before ); assert!(matches!( registry.register_encoder::(|number| { Ok(("replacement".to_owned(), number.0.to_le_bytes().to_vec())) }), Err(CodecRegistrationError::EncoderAlreadyRegistered { .. }) )); assert_eq!( fingerprint(®istry, std::slice::from_ref(&sample)), before ); assert!(matches!( registry.register_decoder::(Number::type_tag(), |_| Ok(Text("replacement".into()))), Err(CodecRegistrationError::DecoderAlreadyRegistered { .. }) )); assert_eq!( fingerprint(®istry, std::slice::from_ref(&sample)), before ); } #[test] fn symmetric_registration_is_atomic_when_only_encoder_conflicts() { let mut registry = CodecRegistry::new(); registry .register_encoder::(|number| { Ok(( Number::type_tag().to_owned(), number.0.to_be_bytes().to_vec(), )) }) .unwrap(); assert!(matches!( registry.register::(NumberCodec), Err(CodecRegistrationError::EncoderAlreadyRegistered { .. }) )); assert!(registry .decode(Number::type_tag(), &0u64.to_be_bytes()) .is_err()); let (tag, bytes) = registry .encode(TypeId::of::(), Box::new(Number(9))) .unwrap(); assert_eq!(tag, Number::type_tag()); assert_eq!(bytes, 9u64.to_be_bytes()); } #[test] fn symmetric_registration_is_atomic_when_only_decoder_conflicts() { let mut registry = CodecRegistry::new(); registry .register_decoder::(Number::type_tag(), |bytes| { Ok(Text(String::from_utf8_lossy(bytes).into_owned())) }) .unwrap(); assert!(matches!( registry.register::(NumberCodec), Err(CodecRegistrationError::DecoderAlreadyRegistered { .. }) )); assert!(registry .encode(TypeId::of::(), Box::new(Number(9))) .is_err()); let decoded = registry.decode(Number::type_tag(), b"first").unwrap(); assert_eq!(*decoded.downcast::().unwrap(), Text("first".into())); } #[test] fn operation_failures_do_not_change_registered_behavior() { let samples = vec![Sample::Number(Number(11)), Sample::Fallible(Fallible(3))]; let mut registry = CodecRegistry::new(); for sample in &samples { sample.register(&mut registry).unwrap(); } let before = fingerprint(®istry, &samples); assert!(registry .encode(TypeId::of::(), Box::new(Text("unknown".into()))) .is_err()); assert_eq!(fingerprint(®istry, &samples), before); assert!(registry .encode(TypeId::of::(), Box::new(Text("wrong".into()))) .is_err()); assert_eq!(fingerprint(®istry, &samples), before); assert!(registry.decode("contract::Unknown", b"anything").is_err()); assert_eq!(fingerprint(®istry, &samples), before); assert!(registry.decode(Number::type_tag(), &[1, 2, 3]).is_err()); assert_eq!(fingerprint(®istry, &samples), before); assert!(registry .encode(TypeId::of::(), Box::new(Fallible(u8::MAX))) .is_err()); assert_eq!(fingerprint(®istry, &samples), before); assert!(registry.decode(Fallible::type_tag(), &[u8::MAX]).is_err()); assert_eq!(fingerprint(®istry, &samples), before); }