gear_subxt/utils/
bits.rs

1// Copyright 2019-2023 Parity Technologies (UK) Ltd.
2// This file is dual-licensed as Apache-2.0 or GPL-3.0.
3// see LICENSE for license details.
4
5//! Generic `scale_bits` over `bitvec`-like `BitOrder` and `BitFormat` types.
6
7use codec::{Compact, Input};
8use scale_bits::{
9    scale::format::{Format, OrderFormat, StoreFormat},
10    Bits,
11};
12use scale_decode::IntoVisitor;
13use std::marker::PhantomData;
14
15/// Associates `bitvec::store::BitStore` trait with corresponding, type-erased `scale_bits::StoreFormat` enum.
16///
17/// Used to decode bit sequences by providing `scale_bits::StoreFormat` using
18/// `bitvec`-like type type parameters.
19pub trait BitStore {
20    /// Corresponding `scale_bits::StoreFormat` value.
21    const FORMAT: StoreFormat;
22    /// Number of bits that the backing store types holds.
23    const BITS: u32;
24}
25macro_rules! impl_store {
26    ($ty:ident, $wrapped:ty) => {
27        impl BitStore for $wrapped {
28            const FORMAT: StoreFormat = StoreFormat::$ty;
29            const BITS: u32 = <$wrapped>::BITS;
30        }
31    };
32}
33impl_store!(U8, u8);
34impl_store!(U16, u16);
35impl_store!(U32, u32);
36impl_store!(U64, u64);
37
38/// Associates `bitvec::order::BitOrder` trait with corresponding, type-erased `scale_bits::OrderFormat` enum.
39///
40/// Used to decode bit sequences in runtime by providing `scale_bits::OrderFormat` using
41/// `bitvec`-like type type parameters.
42pub trait BitOrder {
43    /// Corresponding `scale_bits::OrderFormat` value.
44    const FORMAT: OrderFormat;
45}
46macro_rules! impl_order {
47    ($ty:ident) => {
48        #[doc = concat!("Type-level value that corresponds to `scale_bits::OrderFormat::", stringify!($ty), "` at run-time")]
49        #[doc = concat!(" and `bitvec::order::BitOrder::", stringify!($ty), "` at the type level.")]
50        #[derive(Clone, Debug, PartialEq, Eq)]
51        pub enum $ty {}
52        impl BitOrder for $ty {
53            const FORMAT: OrderFormat = OrderFormat::$ty;
54        }
55    };
56}
57impl_order!(Lsb0);
58impl_order!(Msb0);
59
60/// Constructs a run-time format parameters based on the corresponding type-level parameters.
61fn bit_format<Store: BitStore, Order: BitOrder>() -> Format {
62    Format {
63        order: Order::FORMAT,
64        store: Store::FORMAT,
65    }
66}
67
68/// `scale_bits::Bits` generic over the bit store (`u8`/`u16`/`u32`/`u64`) and bit order (LSB, MSB)
69/// used for SCALE encoding/decoding. Uses `scale_bits::Bits`-default `u8` and LSB format underneath.
70#[derive(Debug, Clone, PartialEq, Eq)]
71pub struct DecodedBits<Store, Order> {
72    bits: Bits,
73    _marker: PhantomData<(Store, Order)>,
74}
75
76impl<Store, Order> DecodedBits<Store, Order> {
77    /// Extracts the underlying `scale_bits::Bits` value.
78    pub fn into_bits(self) -> Bits {
79        self.bits
80    }
81
82    /// References the underlying `scale_bits::Bits` value.
83    pub fn as_bits(&self) -> &Bits {
84        &self.bits
85    }
86}
87
88impl<Store, Order> core::iter::FromIterator<bool> for DecodedBits<Store, Order> {
89    fn from_iter<T: IntoIterator<Item = bool>>(iter: T) -> Self {
90        DecodedBits {
91            bits: Bits::from_iter(iter),
92            _marker: PhantomData,
93        }
94    }
95}
96
97impl<Store: BitStore, Order: BitOrder> codec::Decode for DecodedBits<Store, Order> {
98    fn decode<I: Input>(input: &mut I) -> Result<Self, codec::Error> {
99        /// Equivalent of `BitSlice::MAX_BITS` on 32bit machine.
100        const ARCH32BIT_BITSLICE_MAX_BITS: u32 = 0x1fff_ffff;
101
102        let Compact(bits) = <Compact<u32>>::decode(input)?;
103        // Otherwise it is impossible to store it on 32bit machine.
104        if bits > ARCH32BIT_BITSLICE_MAX_BITS {
105            return Err("Attempt to decode a BitVec with too many bits".into());
106        }
107        // NOTE: Replace with `bits.div_ceil(Store::BITS)` if `int_roundings` is stabilised
108        let elements = (bits / Store::BITS) + u32::from(bits % Store::BITS != 0);
109        let bytes_in_elem = Store::BITS.saturating_div(u8::BITS);
110        let bytes_needed = (elements * bytes_in_elem) as usize;
111
112        // NOTE: We could reduce allocations if it would be possible to directly
113        // decode from an `Input` type using a custom format (rather than default <u8, Lsb0>)
114        // for the `Bits` type.
115        let mut storage = codec::Encode::encode(&Compact(bits));
116        let prefix_len = storage.len();
117        storage.reserve_exact(bytes_needed);
118        storage.extend(vec![0; bytes_needed]);
119        input.read(&mut storage[prefix_len..])?;
120
121        let decoder = scale_bits::decode_using_format_from(&storage, bit_format::<Store, Order>())?;
122        let bits = decoder.collect::<Result<Vec<_>, _>>()?;
123        let bits = Bits::from_iter(bits);
124
125        Ok(DecodedBits {
126            bits,
127            _marker: PhantomData,
128        })
129    }
130}
131
132impl<Store: BitStore, Order: BitOrder> codec::Encode for DecodedBits<Store, Order> {
133    fn size_hint(&self) -> usize {
134        self.bits.size_hint()
135    }
136
137    fn encoded_size(&self) -> usize {
138        self.bits.encoded_size()
139    }
140
141    fn encode(&self) -> Vec<u8> {
142        scale_bits::encode_using_format(self.bits.iter(), bit_format::<Store, Order>())
143    }
144}
145
146#[doc(hidden)]
147pub struct DecodedBitsVisitor<S, O>(std::marker::PhantomData<(S, O)>);
148impl<Store, Order> scale_decode::Visitor for DecodedBitsVisitor<Store, Order> {
149    type Value<'scale, 'info> = DecodedBits<Store, Order>;
150    type Error = scale_decode::Error;
151
152    fn unchecked_decode_as_type<'scale, 'info>(
153        self,
154        input: &mut &'scale [u8],
155        type_id: scale_decode::visitor::TypeId,
156        types: &'info scale_info::PortableRegistry,
157    ) -> scale_decode::visitor::DecodeAsTypeResult<
158        Self,
159        Result<Self::Value<'scale, 'info>, Self::Error>,
160    > {
161        let res = scale_decode::visitor::decode_with_visitor(
162            input,
163            type_id.0,
164            types,
165            Bits::into_visitor(),
166        )
167        .map(|bits| DecodedBits {
168            bits,
169            _marker: PhantomData,
170        });
171        scale_decode::visitor::DecodeAsTypeResult::Decoded(res)
172    }
173}
174impl<Store, Order> scale_decode::IntoVisitor for DecodedBits<Store, Order> {
175    type Visitor = DecodedBitsVisitor<Store, Order>;
176    fn into_visitor() -> Self::Visitor {
177        DecodedBitsVisitor(PhantomData)
178    }
179}
180
181impl<Store, Order> scale_encode::EncodeAsType for DecodedBits<Store, Order> {
182    fn encode_as_type_to(
183        &self,
184        type_id: u32,
185        types: &scale_info::PortableRegistry,
186        out: &mut Vec<u8>,
187    ) -> Result<(), scale_encode::Error> {
188        self.bits.encode_as_type_to(type_id, types, out)
189    }
190}
191
192#[cfg(test)]
193mod tests {
194    use super::*;
195
196    use core::fmt::Debug;
197
198    use bitvec::vec::BitVec;
199    use codec::Decode as _;
200
201    // NOTE: We don't use `bitvec::order` types in our implementation, since we
202    // don't want to depend on `bitvec`. Rather than reimplementing the unsafe
203    // trait on our types here for testing purposes, we simply convert and
204    // delegate to `bitvec`'s own types.
205    trait ToBitVec {
206        type Order: bitvec::order::BitOrder;
207    }
208    impl ToBitVec for Lsb0 {
209        type Order = bitvec::order::Lsb0;
210    }
211    impl ToBitVec for Msb0 {
212        type Order = bitvec::order::Msb0;
213    }
214
215    fn scales_like_bitvec_and_roundtrips<
216        'a,
217        Store: BitStore + bitvec::store::BitStore + PartialEq,
218        Order: BitOrder + ToBitVec + Debug + PartialEq,
219    >(
220        input: impl IntoIterator<Item = &'a bool>,
221    ) where
222        BitVec<Store, <Order as ToBitVec>::Order>: codec::Encode + codec::Decode,
223    {
224        let input: Vec<_> = input.into_iter().copied().collect();
225
226        let decoded_bits = DecodedBits::<Store, Order>::from_iter(input.clone());
227        let bitvec = BitVec::<Store, <Order as ToBitVec>::Order>::from_iter(input);
228
229        let decoded_bits_encoded = codec::Encode::encode(&decoded_bits);
230        let bitvec_encoded = codec::Encode::encode(&bitvec);
231        assert_eq!(decoded_bits_encoded, bitvec_encoded);
232
233        let decoded_bits_decoded =
234            DecodedBits::<Store, Order>::decode(&mut &decoded_bits_encoded[..])
235                .expect("SCALE-encoding DecodedBits to roundtrip");
236        let bitvec_decoded =
237            BitVec::<Store, <Order as ToBitVec>::Order>::decode(&mut &bitvec_encoded[..])
238                .expect("SCALE-encoding BitVec to roundtrip");
239        assert_eq!(decoded_bits, decoded_bits_decoded);
240        assert_eq!(bitvec, bitvec_decoded);
241    }
242
243    #[test]
244    fn decoded_bitvec_scales_and_roundtrips() {
245        let test_cases = [
246            vec![],
247            vec![true],
248            vec![false],
249            vec![true, false, true],
250            vec![true, false, true, false, false, false, false, false, true],
251            [vec![true; 5], vec![false; 5], vec![true; 1], vec![false; 3]].concat(),
252            [vec![true; 9], vec![false; 9], vec![true; 9], vec![false; 9]].concat(),
253        ];
254
255        for test_case in &test_cases {
256            scales_like_bitvec_and_roundtrips::<u8, Lsb0>(test_case);
257            scales_like_bitvec_and_roundtrips::<u16, Lsb0>(test_case);
258            scales_like_bitvec_and_roundtrips::<u32, Lsb0>(test_case);
259            scales_like_bitvec_and_roundtrips::<u64, Lsb0>(test_case);
260            scales_like_bitvec_and_roundtrips::<u8, Msb0>(test_case);
261            scales_like_bitvec_and_roundtrips::<u16, Msb0>(test_case);
262            scales_like_bitvec_and_roundtrips::<u32, Msb0>(test_case);
263            scales_like_bitvec_and_roundtrips::<u64, Msb0>(test_case);
264        }
265    }
266}