Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
91 changes: 90 additions & 1 deletion crates/sats/src/bsatn.rs
Original file line number Diff line number Diff line change
Expand Up @@ -231,10 +231,99 @@ mod tests {
use super::{to_vec, DecodeError, Deserializer};
use crate::de::DeserializeSeed;
use crate::proptest::{generate_algebraic_type, generate_typed_value};
use crate::{meta_type::MetaType, AlgebraicType, AlgebraicValue, WithTypespace};
use crate::{
meta_type::MetaType, AlgebraicType, AlgebraicTypeRef, AlgebraicValue, ArrayValue, Typespace, WithTypespace,
};
use proptest::prelude::*;
use proptest::{collection::vec, proptest};

#[test]
fn decode_invalid_type_reference() {
let ty: AlgebraicType = super::from_slice(&[0, 0, 0, 0, 0]).unwrap();
assert_eq!(
AlgebraicValue::decode(&ty, &mut &[0u8; 8][..]),
Err(DecodeError::Other("Type reference &0 out of bounds".into()))
);
}

fn check_invalid_reference(typespace: &Typespace, ty: &AlgebraicType, bytes: &[u8], invalid: u32) {
let seed = WithTypespace::new(typespace, ty);
let expected = DecodeError::Other(format!("Type reference &{invalid} out of bounds"));
assert_eq!(seed.deserialize(Deserializer::new(&mut &*bytes)), Err(expected.clone()));
assert_eq!(seed.validate(Deserializer::new(&mut &*bytes)), Err(expected.clone()));
if let AlgebraicType::Array(array) = ty {
let seed = seed.with(array);
assert_eq!(seed.deserialize(Deserializer::new(&mut &*bytes)), Err(expected.clone()));
assert_eq!(seed.validate(Deserializer::new(&mut &*bytes)), Err(expected));
}
}

#[test]
fn decode_and_validate_invalid_type_references() {
for typespace in [Typespace::default(), Typespace::new(vec![AlgebraicType::U8])] {
for invalid in [typespace.types.len() as u32, u32::MAX] {
let ty = AlgebraicType::Ref(AlgebraicTypeRef(invalid));
check_invalid_reference(&typespace, &ty, &[], invalid);
check_invalid_reference(&typespace, &AlgebraicType::product([ty.clone()]), &[], invalid);
check_invalid_reference(&typespace, &AlgebraicType::option(ty.clone()), &[0], invalid);
for bytes in [&[0, 0, 0, 0][..], &[1, 0, 0, 0, 42][..]] {
check_invalid_reference(&typespace, &AlgebraicType::array(ty.clone()), bytes, invalid);
}
}
}

let ty = AlgebraicType::Ref(AlgebraicTypeRef(0));
let typespace = Typespace::new(vec![AlgebraicType::Ref(AlgebraicTypeRef(1))]);
check_invalid_reference(&typespace, &ty, &[], 1);
check_invalid_reference(&typespace, &AlgebraicType::array(ty), &[0, 0, 0, 0], 1);

// An unselected sum branch does not need to be resolved.
check_reference_value(
Typespace::EMPTY,
&AlgebraicType::option(AlgebraicType::Ref(AlgebraicTypeRef(0))),
AlgebraicValue::sum(1, AlgebraicValue::unit()),
);
}

fn check_reference_value(typespace: &Typespace, ty: &AlgebraicType, value: AlgebraicValue) {
let bytes = to_vec(&value).unwrap();
let seed = WithTypespace::new(typespace, ty);
let mut input = &bytes[..];
assert_eq!(seed.deserialize(Deserializer::new(&mut input)), Ok(value));
assert!(input.is_empty());
let mut input = &bytes[..];
assert_eq!(seed.validate(Deserializer::new(&mut input)), Ok(()));
assert!(input.is_empty());
}

#[test]
fn decode_and_validate_valid_type_references() {
let ty = AlgebraicType::Ref(AlgebraicTypeRef(0));
let typespace = Typespace::new(vec![AlgebraicType::Ref(AlgebraicTypeRef(1)), AlgebraicType::U8]);
check_reference_value(&typespace, &ty, AlgebraicValue::U8(42));
check_reference_value(
&typespace,
&AlgebraicType::array(ty),
AlgebraicValue::Array(ArrayValue::U8(vec![42, 7].into())),
);
}

#[test]
fn decode_and_validate_recursive_type_references() {
let ty = AlgebraicType::Ref(AlgebraicTypeRef(0));
let typespace = Typespace::new(vec![AlgebraicType::option(ty.clone())]);
let nil = AlgebraicValue::sum(1, AlgebraicValue::unit());
check_reference_value(&typespace, &ty, AlgebraicValue::sum(0, AlgebraicValue::sum(0, nil)));

let typespace = Typespace::new(vec![AlgebraicType::array(ty.clone())]);
let empty = ArrayValue::Array(vec![].into());
check_reference_value(
&typespace,
&ty,
AlgebraicValue::Array(ArrayValue::Array(vec![empty].into())),
);
}

#[test]
fn type_to_binary_equivalent() {
check_type(&AlgebraicType::meta_type());
Expand Down
19 changes: 15 additions & 4 deletions crates/sats/src/de/impls.rs
Original file line number Diff line number Diff line change
Expand Up @@ -8,6 +8,7 @@ use crate::{
SumValue, WithTypespace, F32, F64,
};
use crate::{i256, u256};
use crate::{typespace::TypeRefError, AlgebraicTypeRef, Typespace};
use core::{iter, marker::PhantomData, ops::Bound};
use lean_string::LeanString;
use smallvec::SmallVec;
Expand Down Expand Up @@ -580,12 +581,20 @@ impl<'de, T: Copy + DeserializeSeed<'de>> VariantVisitor<'de> for BoundVisitor<T
}
}

fn resolve_type_ref<E: Error>(typespace: &Typespace, r: AlgebraicTypeRef) -> Result<&AlgebraicType, E> {
typespace
.get(r)
.ok_or_else(|| E::custom(TypeRefError::InvalidTypeRef(r)))
}

impl<'de> DeserializeSeed<'de> for WithTypespace<'_, AlgebraicType> {
type Output = AlgebraicValue;

fn deserialize<D: Deserializer<'de>>(self, de: D) -> Result<Self::Output, D::Error> {
match self.ty() {
AlgebraicType::Ref(r) => self.resolve(*r).deserialize(de),
AlgebraicType::Ref(r) => self
.with(resolve_type_ref::<D::Error>(self.typespace(), *r)?)
.deserialize(de),
AlgebraicType::Sum(sum) => self.with(sum).deserialize(de).map(Into::into),
AlgebraicType::Product(prod) => self.with(prod).deserialize(de).map(Into::into),
AlgebraicType::Array(ty) => self.with(ty).deserialize(de).map(Into::into),
Expand All @@ -610,7 +619,9 @@ impl<'de> DeserializeSeed<'de> for WithTypespace<'_, AlgebraicType> {

fn validate<D: Deserializer<'de>>(self, de: D) -> Result<(), D::Error> {
match self.ty() {
AlgebraicType::Ref(r) => self.resolve(*r).validate(de),
AlgebraicType::Ref(r) => self
.with(resolve_type_ref::<D::Error>(self.typespace(), *r)?)
.validate(de),
AlgebraicType::Sum(sum) => self.with(sum).validate(de),
AlgebraicType::Product(prod) => self.with(prod).validate(de),
AlgebraicType::Array(ty) => self.with(ty).validate(de),
Expand Down Expand Up @@ -774,7 +785,7 @@ impl<'de> DeserializeSeed<'de> for WithTypespace<'_, ArrayType> {
break match ty {
AlgebraicType::Ref(r) => {
// The only arm that will loop.
ty = self.resolve(*r).ty();
ty = resolve_type_ref::<D::Error>(self.typespace(), *r)?;
continue;
}
AlgebraicType::Sum(ty) => deserializer
Expand Down Expand Up @@ -825,7 +836,7 @@ impl<'de> DeserializeSeed<'de> for WithTypespace<'_, ArrayType> {
break match ty {
AlgebraicType::Ref(r) => {
// The only arm that will loop.
ty = self.resolve(*r).ty();
ty = resolve_type_ref::<D::Error>(self.typespace(), *r)?;
continue;
}
AlgebraicType::Sum(ty) => deserializer.validate_array_seed(BasicVecVisitor, self.with(ty)),
Expand Down