1use codec::{Compact, Input};
8use scale_bits::{
9 scale::format::{Format, OrderFormat, StoreFormat},
10 Bits,
11};
12use scale_decode::IntoVisitor;
13use std::marker::PhantomData;
14
15pub trait BitStore {
20 const FORMAT: StoreFormat;
22 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
38pub trait BitOrder {
43 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
60fn bit_format<Store: BitStore, Order: BitOrder>() -> Format {
62 Format {
63 order: Order::FORMAT,
64 store: Store::FORMAT,
65 }
66}
67
68#[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 pub fn into_bits(self) -> Bits {
79 self.bits
80 }
81
82 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 const ARCH32BIT_BITSLICE_MAX_BITS: u32 = 0x1fff_ffff;
101
102 let Compact(bits) = <Compact<u32>>::decode(input)?;
103 if bits > ARCH32BIT_BITSLICE_MAX_BITS {
105 return Err("Attempt to decode a BitVec with too many bits".into());
106 }
107 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 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 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}