Skip to main content

vstd/
bits.rs

1//! Properties of bitwise operators.
2use super::prelude::*;
3
4verus! {
5
6#[cfg(verus_keep_ghost)]
7use super::arithmetic::power::pow;
8#[cfg(verus_keep_ghost)]
9use super::arithmetic::power2::{
10    pow2,
11    lemma_pow2_unfold,
12    lemma_pow2_adds,
13    lemma_pow2_pos,
14    lemma2_to64,
15    lemma2_to64_rest,
16    lemma_pow2_strictly_increases,
17};
18#[cfg(verus_keep_ghost)]
19use super::arithmetic::div_mod::{
20    lemma_div_by_multiple,
21    lemma_div_denominator,
22    lemma_div_is_ordered,
23    lemma_mod_breakdown,
24    lemma_mod_multiples_vanish,
25    lemma_remainder_lower,
26};
27#[cfg(verus_keep_ghost)]
28use super::arithmetic::mul::{
29    lemma_mul_inequality,
30    lemma_mul_is_commutative,
31    lemma_mul_is_associative,
32};
33#[cfg(verus_keep_ghost)]
34use super::calc_macro::*;
35
36} // verus!
37// Proofs that shift right is equivalent to division by power of 2.
38macro_rules! lemma_shr_is_div {
39    ($name:ident, $uN:ty) => {
40        #[cfg(verus_keep_ghost)]
41        verus! {
42        #[doc = "Proof that for x and n of type "]
43        #[doc = stringify!($uN)]
44        #[doc = ", shifting x right by n is equivalent to division of x by 2^n."]
45        pub broadcast proof fn $name(x: $uN, shift: $uN)
46            requires
47                0 <= shift < <$uN>::BITS,
48            ensures
49                #[trigger] (x >> shift) == x as nat / pow2(shift as nat),
50            decreases shift,
51        {
52            // Step by 4 to reduce recursion depth (divisor 16 fits in all unsigned types).
53            reveal(pow);
54            if shift == 0 {
55                assert(x >> 0 == x) by (bit_vector);
56                assert(pow2(0) == 1) by (compute_only);
57            } else if shift == 1 {
58                assert(x >> 1 == x / 2) by (bit_vector);
59                assert(pow2(1) == 2) by (compute_only);
60            } else if shift == 2 {
61                assert(x >> 2 == x / 4) by (bit_vector);
62                assert(pow2(2) == 4) by (compute_only);
63            } else if shift == 3 {
64                assert(x >> 3 == x / 8) by (bit_vector);
65                assert(pow2(3) == 8) by (compute_only);
66            } else {
67                assert(x >> shift == (x >> (sub(shift, 4) as $uN)) / 16) by (bit_vector)
68                    requires
69                        4 <= shift < <$uN>::BITS,
70                ;
71                calc!{ (==)
72                    (x >> shift) as nat;
73                        {}
74                    ((x >> (sub(shift, 4) as $uN)) / 16) as nat;
75                        { $name(x, (shift - 4) as $uN); }
76                    (x as nat / pow2((shift - 4) as nat)) / 16;
77                        {
78                            lemma_pow2_pos((shift - 4) as nat);
79                            lemma2_to64();
80                            assert(pow2(4) == 16) by (compute_only);
81                            lemma_div_denominator(x as int, pow2((shift - 4) as nat) as int, 16);
82                        }
83                    x as nat / (pow2((shift - 4) as nat) * pow2(4));
84                        {
85                            lemma_pow2_adds((shift - 4) as nat, 4);
86                        }
87                    x as nat / pow2(shift as nat);
88                }
89            }
90        }
91        }
92    };
93}
94
95lemma_shr_is_div!(lemma_u128_shr_is_div, u128);
96lemma_shr_is_div!(lemma_u64_shr_is_div, u64);
97lemma_shr_is_div!(lemma_u32_shr_is_div, u32);
98lemma_shr_is_div!(lemma_u16_shr_is_div, u16);
99lemma_shr_is_div!(lemma_u8_shr_is_div, u8);
100lemma_shr_is_div!(lemma_usize_shr_is_div, usize);
101
102// Proofs of when a power of 2 fits in an unsigned type.
103macro_rules! lemma_pow2_no_overflow {
104    ($name:ident, $uN:ty) => {
105        #[cfg(verus_keep_ghost)]
106        verus! {
107        #[doc = "Proof that 2^n does not overflow "]
108        #[doc = stringify!($uN)]
109        #[doc = " for an exponent n."]
110        pub broadcast proof fn $name(n: nat)
111            requires
112                0 <= n < <$uN>::BITS,
113            ensures
114                0 < #[trigger] pow2(n) < <$uN>::MAX,
115        {
116            lemma_pow2_pos(n);
117            lemma2_to64();
118            lemma2_to64_rest();
119        }
120        }
121    };
122}
123
124lemma_pow2_no_overflow!(lemma_u64_pow2_no_overflow, u64);
125lemma_pow2_no_overflow!(lemma_u32_pow2_no_overflow, u32);
126lemma_pow2_no_overflow!(lemma_u16_pow2_no_overflow, u16);
127lemma_pow2_no_overflow!(lemma_u8_pow2_no_overflow, u8);
128lemma_pow2_no_overflow!(lemma_usize_pow2_no_overflow, usize);
129
130// Proofs that shift left is equivalent to multiplication by power of 2.
131macro_rules! lemma_shl_is_mul {
132    ($name:ident, $no_overflow:ident, $uN:ty) => {
133        #[cfg(verus_keep_ghost)]
134        verus! {
135        #[doc = "Proof that for x and n of type "]
136        #[doc = stringify!($uN)]
137        #[doc = ", shifting x left by n is equivalent to multiplication of x by 2^n (provided no overflow)."]
138        pub broadcast proof fn $name(x: $uN, shift: $uN)
139            requires
140                0 <= shift < <$uN>::BITS,
141                x * pow2(shift as nat) <= <$uN>::MAX,
142            ensures
143                #[trigger] (x << shift) == x * pow2(shift as nat),
144            decreases shift,
145        {
146            $no_overflow(shift as nat);
147            if shift == 0 {
148                assert(x << 0 == x) by (bit_vector);
149                assert(pow2(0) == 1) by (compute_only);
150                super::arithmetic::mul::lemma_mul_basics(x as int);
151                assert((x << shift) == x * pow2(shift as nat));
152            } else {
153                assert(x << shift == mul(x << ((sub(shift, 1)) as $uN), 2)) by (bit_vector)
154                    requires
155                        0 < shift < <$uN>::BITS,
156                ;
157                assert((x << (sub(shift, 1) as $uN)) == x * pow2(sub(shift, 1) as nat)) by {
158                    lemma_pow2_strictly_increases((shift - 1) as nat, shift as nat);
159                    lemma_mul_inequality(
160                        pow2((shift - 1) as nat) as int,
161                        pow2(shift as nat) as int,
162                        x as int,
163                    );
164                    lemma_mul_is_commutative(x as int, pow2((shift - 1) as nat) as int);
165                    lemma_mul_is_commutative(x as int, pow2(shift as nat) as int);
166                    $name(x, (shift - 1) as $uN);
167                }
168                calc!{ (==)
169                    ((x << (sub(shift, 1) as $uN)) * 2);
170                        {}
171                    ((x * pow2(sub(shift, 1) as nat)) * 2);
172                        {
173                            lemma_mul_is_associative(x as int, pow2(sub(shift, 1) as nat) as int, 2);
174                        }
175                    x * ((pow2(sub(shift, 1) as nat)) * 2);
176                        {
177                            lemma_pow2_adds((shift - 1) as nat, 1);
178                            lemma2_to64();
179                        }
180                    x * pow2(shift as nat);
181                }
182                assert((x << shift) == x * pow2(shift as nat));
183            }
184        }
185        }
186    };
187}
188
189lemma_shl_is_mul!(lemma_u64_shl_is_mul, lemma_u64_pow2_no_overflow, u64);
190lemma_shl_is_mul!(lemma_u32_shl_is_mul, lemma_u32_pow2_no_overflow, u32);
191lemma_shl_is_mul!(lemma_u16_shl_is_mul, lemma_u16_pow2_no_overflow, u16);
192lemma_shl_is_mul!(lemma_u8_shl_is_mul, lemma_u8_pow2_no_overflow, u8);
193lemma_shl_is_mul!(lemma_usize_shl_is_mul, lemma_usize_pow2_no_overflow, usize);
194
195macro_rules! lemma_mul_pow2_le_max_iff_max_shr {
196    ($name:ident, $shr_is_div:ident, $uN:ty) => {
197        #[cfg(verus_keep_ghost)]
198        verus! {
199        #[doc = "Proof that for x, n and max of type "]
200        #[doc = stringify!($uN)]
201        #[doc = ", multiplication of x by 2^n is less than or equal to max if and only if x is less than or equal to shifting max right by n."]
202        pub proof fn $name(x: $uN, shift: $uN, max: $uN)
203        requires
204            0 <= shift < <$uN>::BITS,
205        ensures
206            x * pow2(shift as nat) <= max <==> x <= (max >> shift),
207    {
208        assert(max >> shift == max as nat / pow2(shift as nat)) by {
209            $shr_is_div(max, shift as $uN);
210        };
211
212        lemma_pow2_pos(shift as nat);
213
214        if x * pow2(shift as nat) <= max {
215            assert(x <= (max as nat) / pow2(shift as nat)) by {
216                lemma_div_is_ordered(x as int * pow2(shift as nat) as int, max as int, pow2(shift as nat) as int);
217                lemma_div_by_multiple(x as int, pow2(shift as nat) as int);
218            };
219        }
220        if x <= (max >> shift) {
221            assert(x * pow2(shift as nat) <= max as nat) by {
222                lemma_mul_inequality(x as int, max as int / pow2(shift as nat) as int,  pow2(shift as nat) as int);
223                lemma_remainder_lower(max as int, pow2(shift as nat) as int);
224                lemma_mul_is_commutative(max as int / pow2(shift as nat) as int,  pow2(shift as nat) as int);
225            };
226        }
227    }
228    }
229    };
230}
231
232lemma_mul_pow2_le_max_iff_max_shr!(
233    lemma_u64_mul_pow2_le_max_iff_max_shr,
234    lemma_u64_shr_is_div,
235    u64
236);
237lemma_mul_pow2_le_max_iff_max_shr!(
238    lemma_u32_mul_pow2_le_max_iff_max_shr,
239    lemma_u32_shr_is_div,
240    u32
241);
242lemma_mul_pow2_le_max_iff_max_shr!(
243    lemma_u16_mul_pow2_le_max_iff_max_shr,
244    lemma_u16_shr_is_div,
245    u16
246);
247lemma_mul_pow2_le_max_iff_max_shr!(lemma_u8_mul_pow2_le_max_iff_max_shr, lemma_u8_shr_is_div, u8);
248lemma_mul_pow2_le_max_iff_max_shr!(
249    lemma_usize_mul_pow2_le_max_iff_max_shr,
250    lemma_usize_shr_is_div,
251    usize
252);
253
254verus! {
255
256/// Mask with low n bits set.
257pub open spec fn low_bits_mask(n: nat) -> nat {
258    (pow2(n) - 1) as nat
259}
260
261/// Proof relating the n-bit mask to a function of the (n-1)-bit mask.
262pub broadcast proof fn lemma_low_bits_mask_unfold(n: nat)
263    requires
264        n > 0,
265    ensures
266        #[trigger] low_bits_mask(n) == 2 * low_bits_mask((n - 1) as nat) + 1,
267{
268    calc! {
269        (==)
270        low_bits_mask(n); {}
271        (pow2(n) - 1) as nat; {
272            lemma_pow2_unfold(n);
273        }
274        (2 * pow2((n - 1) as nat) - 1) as nat; {}
275        (2 * (pow2((n - 1) as nat) - 1) + 1) as nat; {
276            lemma_pow2_pos((n - 1) as nat);
277        }
278        (2 * low_bits_mask((n - 1) as nat) + 1) as nat;
279    }
280}
281
282/// Proof that low_bits_mask(n) is odd.
283pub broadcast proof fn lemma_low_bits_mask_is_odd(n: nat)
284    requires
285        n > 0,
286    ensures
287        #[trigger] (low_bits_mask(n) % 2) == 1,
288{
289    calc! {
290        (==)
291        low_bits_mask(n) % 2; {
292            lemma_low_bits_mask_unfold(n);
293        }
294        (2 * low_bits_mask((n - 1) as nat) + 1) % 2; {
295            lemma_mod_multiples_vanish(low_bits_mask((n - 1) as nat) as int, 1, 2);
296        }
297        1nat % 2;
298    }
299}
300
301/// Proof that dividing the low n bit mask by 2 gives the low n-1 bit mask.
302pub broadcast proof fn lemma_low_bits_mask_div2(n: nat)
303    requires
304        n > 0,
305    ensures
306        #[trigger] (low_bits_mask(n) / 2) == low_bits_mask((n - 1) as nat),
307{
308    lemma_low_bits_mask_unfold(n);
309}
310
311/// Proof establishing the concrete values of all masks of bit sizes from 0 to
312/// 32, and 64.
313pub proof fn lemma_low_bits_mask_values()
314    ensures
315        low_bits_mask(0) == 0x0,
316        low_bits_mask(1) == 0x1,
317        low_bits_mask(2) == 0x3,
318        low_bits_mask(3) == 0x7,
319        low_bits_mask(4) == 0xf,
320        low_bits_mask(5) == 0x1f,
321        low_bits_mask(6) == 0x3f,
322        low_bits_mask(7) == 0x7f,
323        low_bits_mask(8) == 0xff,
324        low_bits_mask(9) == 0x1ff,
325        low_bits_mask(10) == 0x3ff,
326        low_bits_mask(11) == 0x7ff,
327        low_bits_mask(12) == 0xfff,
328        low_bits_mask(13) == 0x1fff,
329        low_bits_mask(14) == 0x3fff,
330        low_bits_mask(15) == 0x7fff,
331        low_bits_mask(16) == 0xffff,
332        low_bits_mask(17) == 0x1ffff,
333        low_bits_mask(18) == 0x3ffff,
334        low_bits_mask(19) == 0x7ffff,
335        low_bits_mask(20) == 0xfffff,
336        low_bits_mask(21) == 0x1fffff,
337        low_bits_mask(22) == 0x3fffff,
338        low_bits_mask(23) == 0x7fffff,
339        low_bits_mask(24) == 0xffffff,
340        low_bits_mask(25) == 0x1ffffff,
341        low_bits_mask(26) == 0x3ffffff,
342        low_bits_mask(27) == 0x7ffffff,
343        low_bits_mask(28) == 0xfffffff,
344        low_bits_mask(29) == 0x1fffffff,
345        low_bits_mask(30) == 0x3fffffff,
346        low_bits_mask(31) == 0x7fffffff,
347        low_bits_mask(32) == 0xffffffff,
348        low_bits_mask(64) == 0xffffffffffffffff,
349{
350    #[verusfmt::skip]
351    assert(
352        low_bits_mask(0) == 0x0 &&
353        low_bits_mask(1) == 0x1 &&
354        low_bits_mask(2) == 0x3 &&
355        low_bits_mask(3) == 0x7 &&
356        low_bits_mask(4) == 0xf &&
357        low_bits_mask(5) == 0x1f &&
358        low_bits_mask(6) == 0x3f &&
359        low_bits_mask(7) == 0x7f &&
360        low_bits_mask(8) == 0xff &&
361        low_bits_mask(9) == 0x1ff &&
362        low_bits_mask(10) == 0x3ff &&
363        low_bits_mask(11) == 0x7ff &&
364        low_bits_mask(12) == 0xfff &&
365        low_bits_mask(13) == 0x1fff &&
366        low_bits_mask(14) == 0x3fff &&
367        low_bits_mask(15) == 0x7fff &&
368        low_bits_mask(16) == 0xffff &&
369        low_bits_mask(17) == 0x1ffff &&
370        low_bits_mask(18) == 0x3ffff &&
371        low_bits_mask(19) == 0x7ffff &&
372        low_bits_mask(20) == 0xfffff &&
373        low_bits_mask(21) == 0x1fffff &&
374        low_bits_mask(22) == 0x3fffff &&
375        low_bits_mask(23) == 0x7fffff &&
376        low_bits_mask(24) == 0xffffff &&
377        low_bits_mask(25) == 0x1ffffff &&
378        low_bits_mask(26) == 0x3ffffff &&
379        low_bits_mask(27) == 0x7ffffff &&
380        low_bits_mask(28) == 0xfffffff &&
381        low_bits_mask(29) == 0x1fffffff &&
382        low_bits_mask(30) == 0x3fffffff &&
383        low_bits_mask(31) == 0x7fffffff &&
384        low_bits_mask(32) == 0xffffffff &&
385        low_bits_mask(64) == 0xffffffffffffffff
386    ) by (compute_only);
387}
388
389} // verus!
390// Proofs that and with mask is equivalent to modulo with power of two.
391macro_rules! lemma_low_bits_mask_is_mod {
392    ($name:ident, $and_split_low_bit:ident, $no_overflow:ident, $uN:ty) => {
393        #[cfg(verus_keep_ghost)]
394        verus! {
395        #[doc = "Proof that for natural n and x of type "]
396        #[doc = stringify!($uN)]
397        #[doc = ", and with the low n-bit mask is equivalent to modulo 2^n."]
398        pub broadcast proof fn $name(x: $uN, n: nat)
399            requires
400                n < <$uN>::BITS,
401            ensures
402                #[trigger] (x & (low_bits_mask(n) as $uN)) == x % (pow2(n) as $uN),
403            decreases n,
404        {
405            // Bounds.
406            $no_overflow(n);
407            lemma_pow2_pos(n);
408
409            // Inductive proof.
410            if n == 0 {
411                assert(low_bits_mask(0) == 0) by (compute_only);
412                assert(x & 0 == 0) by (bit_vector);
413                assert(pow2(0) == 1) by (compute_only);
414                assert(x % 1 == 0);
415            } else {
416                lemma_pow2_unfold(n);
417                assert((x % 2) == ((x % 2) & 1)) by (bit_vector);
418                calc!{ (==)
419                    x % (pow2(n) as $uN);
420                        {}
421                    x % ((2 * pow2((n-1) as nat)) as $uN);
422                        {
423                            lemma_pow2_pos((n-1) as nat);
424                            lemma_mod_breakdown(x as int, 2, pow2((n-1) as nat) as int);
425                        }
426                    add(mul(2, (x / 2) % (pow2((n-1) as nat) as $uN)), x % 2);
427                        {
428                            $name(x/2, (n-1) as nat);
429                        }
430                    add(mul(2, (x / 2) & (low_bits_mask((n-1) as nat) as $uN)), x % 2);
431                        {
432                            lemma_low_bits_mask_div2(n);
433                        }
434                    add(mul(2, (x / 2) & (low_bits_mask(n) as $uN / 2)), x % 2);
435                        {
436                            lemma_low_bits_mask_is_odd(n);
437                        }
438                    add(mul(2, (x / 2) & (low_bits_mask(n) as $uN / 2)), (x % 2) & ((low_bits_mask(n) as $uN) % 2));
439                        {
440                            $and_split_low_bit(x as $uN, low_bits_mask(n) as $uN);
441                        }
442                    x & (low_bits_mask(n) as $uN);
443                }
444            }
445        }
446
447        // Helper lemma breaking a bitwise-and operation into the low bit and the rest.
448        proof fn $and_split_low_bit(x: $uN, m: $uN)
449            by (bit_vector)
450            ensures
451                x & m == add(mul(((x / 2) & (m / 2)), 2), (x % 2) & (m % 2)),
452        {
453        }
454        }
455    };
456}
457
458lemma_low_bits_mask_is_mod!(
459    lemma_u64_low_bits_mask_is_mod,
460    lemma_u64_and_split_low_bit,
461    lemma_u64_pow2_no_overflow,
462    u64
463);
464lemma_low_bits_mask_is_mod!(
465    lemma_u32_low_bits_mask_is_mod,
466    lemma_u32_and_split_low_bit,
467    lemma_u32_pow2_no_overflow,
468    u32
469);
470lemma_low_bits_mask_is_mod!(
471    lemma_u16_low_bits_mask_is_mod,
472    lemma_u16_and_split_low_bit,
473    lemma_u16_pow2_no_overflow,
474    u16
475);
476lemma_low_bits_mask_is_mod!(
477    lemma_u8_low_bits_mask_is_mod,
478    lemma_u8_and_split_low_bit,
479    lemma_u8_pow2_no_overflow,
480    u8
481);
482lemma_low_bits_mask_is_mod!(
483    lemma_usize_low_bits_mask_is_mod,
484    lemma_usize_and_split_low_bit,
485    lemma_usize_pow2_no_overflow,
486    usize
487);