1use 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} macro_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 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
102macro_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
130macro_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
256pub open spec fn low_bits_mask(n: nat) -> nat {
258 (pow2(n) - 1) as nat
259}
260
261pub 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
282pub 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
301pub 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
311pub 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} macro_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 $no_overflow(n);
407 lemma_pow2_pos(n);
408
409 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 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);