1use core::{marker::PhantomData, ptr::NonNull};
3
4use core::num::NonZeroUsize;
5use vstd::raw_ptr::group_raw_ptr_axioms;
6use vstd::{
7 bits,
8 prelude::*,
9 std_specs::{nonzero::*, ops::BitOrSpec},
10};
11use vstd_extra::{prelude::*, sum::Sum};
12
13use super::{NonNullPtr, NonNullPtrRef};
14use crate::util::Either;
15
16verus! {
17
18broadcast use {group_nonull_axioms, group_nonzero_axioms, group_raw_ptr_axioms};
19unsafe impl<L: NonNullPtr, R: NonNullPtr> NonNullPtr for Either<L, R> {
24 type Target = PhantomData<Self>;
25
26 type Permission = Sum<L::Permission, R::Permission>;
31
32 #[verifier::external_body]
33 const ALIGN_BITS: u32 = min(L::ALIGN_BITS, R::ALIGN_BITS).checked_sub(1).expect(
34 "`L` and `R` alignments should be at least 2 to pack `Either` into one pointer",
35 );
36
37 #[verifier::spinoff_prover]
38 fn into_raw(self) -> (ret: (NonNull<Self::Target>, Tracked<Self::Permission>)) {
39 proof_decl!{
40 let ghost align_bits = Self::ALIGN_BITS;
41 let ghost l_align_bits = L::ALIGN_BITS;
42 let ghost r_align_bits = R::ALIGN_BITS;
43 let ghost tag = 1usize << align_bits;
44 }
45 proof! {
46 L::lemma_align_bits_range();
47 R::lemma_align_bits_range();
48 Self::lemma_align_bits_range();
49 vstd::bits::lemma_usize_pow2_no_overflow(align_bits as nat);
50 vstd::bits::lemma_usize_pow2_no_overflow(l_align_bits as nat);
51 vstd::bits::lemma_usize_pow2_no_overflow(r_align_bits as nat);
52 vstd::bits::lemma_usize_shl_is_mul(1, align_bits as usize);
53 vstd::bits::lemma_usize_shl_is_mul(1, l_align_bits as usize);
54 vstd::bits::lemma_usize_shl_is_mul(1, r_align_bits as usize);
55 }
56 match self {
57 Self::Left(left) => {
58 let (left, Tracked(perm)) = left.into_raw();
60 proof! {
61 let left_addr = left.cast::<Self::Target>().view_ptr_mut().addr();
62 let extra_bits: u32 = (l_align_bits - align_bits) as u32;
63 let scale = 1usize << extra_bits;
64 vstd::bits::lemma_usize_pow2_no_overflow(extra_bits as nat);
65 vstd::bits::lemma_usize_shl_is_mul(1, extra_bits as usize);
66 vstd::arithmetic::power2::lemma_pow2_adds(align_bits as nat, extra_bits as nat);
67 assert(left_addr % tag == 0) by {
68 let big = 1usize << l_align_bits;
69 let q = left_addr / big;
70 vstd::arithmetic::div_mod::lemma_fundamental_div_mod(left_addr as int, big as int);
71 assert(left_addr == q * scale * tag) by (nonlinear_arith)
72 requires
73 left_addr == q * big,
74 big == tag * scale,
75 ;
76 vstd::arithmetic::div_mod::lemma_mod_multiples_basic(q * scale, tag as int);
77 };
78 lemma_aligned_addr_clears_tag_bit(left_addr, tag, align_bits, l_align_bits);
79 }
80 (left.cast(), Tracked(Sum::Left(perm)))
81 },
82 Self::Right(right) => {
83 let (right, Tracked(perm)) = right.into_raw();
88 let right_tagged = right.map_addr_v(
89 |addr: NonZeroUsize| -> (ret: NonZeroUsize)
90 ensures
91 ret@ == addr@ | (1usize << Self::ALIGN_BITS),
92 {
93 proof {
94 let tag = 1usize << Self::ALIGN_BITS;
95 let a = addr@;
96 assert(a | tag != 0) by (bit_vector)
97 requires
98 a != 0,
99 ;
100 }
101 addr | 1usize << Self::ALIGN_BITS
102 },
103 );
104 proof! {
105 let addr = right.addr_spec()@;
106 let tagged_addr = right_tagged.addr_spec()@;
107 assert(tagged_addr & tag == tag) by (bit_vector)
108 requires
109 tagged_addr == addr | tag,
110 tag == 1usize << align_bits,
111 1 <= r_align_bits < usize::BITS,
112 align_bits < r_align_bits,
113 addr % (1usize << r_align_bits) == 0,
114 addr != 0,
115 ;
116 let extra_bits: u32 = (r_align_bits - align_bits) as u32;
117 let scale = 1usize << extra_bits;
118 vstd::bits::lemma_usize_pow2_no_overflow(extra_bits as nat);
119 vstd::bits::lemma_usize_shl_is_mul(1, extra_bits as usize);
120 vstd::arithmetic::power2::lemma_pow2_adds(align_bits as nat, extra_bits as nat);
121 assert(tagged_addr == addr + tag) by (bit_vector)
122 requires
123 tagged_addr == addr | tag,
124 tag == 1usize << align_bits,
125 1 <= r_align_bits < usize::BITS,
126 align_bits < r_align_bits,
127 addr % (1usize << r_align_bits) == 0,
128 addr != 0,
129 ;
130 assert(addr % tag == 0) by {
131 let big = 1usize << r_align_bits;
132 let q = addr / big;
133 vstd::arithmetic::div_mod::lemma_fundamental_div_mod(addr as int, big as int);
134 assert(addr == q * scale * tag) by (nonlinear_arith)
135 requires
136 addr == q * big,
137 big == tag * scale,
138 ;
139 vstd::arithmetic::div_mod::lemma_mod_multiples_basic(q * scale, tag as int);
140 }
141 lemma_aligned_addr_clears_tag_bit(addr, tag, align_bits, r_align_bits);
142 assert(tagged_addr & !tag == addr) by (bit_vector)
143 requires
144 tagged_addr == addr | tag,
145 addr & tag == 0,
146 ;
147 assert(tagged_addr % (1usize << align_bits) == 0) by {
148 vstd::arithmetic::div_mod::lemma_mod_add_multiples_vanish(addr as int, tag as int);
149 }
150 }
151 (right_tagged.cast(), Tracked(Sum::Right(perm)))
152 },
153 }
154 }
155
156 unsafe fn from_raw(
157 ptr: NonNull<Self::Target>,
158 Tracked(perm): Tracked<Self::Permission>,
159 ) -> Self {
160 proof! {
161 Self::lemma_align_bits_range();
162 }
163 proof_decl! {
164 let ghost align_bits = Self::ALIGN_BITS;
165 let ghost tag = 1usize << Self::ALIGN_BITS;
166 let ghost ptr_addr = ptr.view_ptr_mut()@.addr;
167 }
168 proof! {
169 assert(tag > 0) by (bit_vector)
170 requires
171 tag == 1usize << align_bits,
172 align_bits < usize::BITS,
173 ;
174 match perm {
175 Sum::Left(_) => {
176 assert((ptr_addr & !tag) == ptr_addr) by (bit_vector)
177 requires
178 ptr_addr & tag == 0,
179 ;
180 },
181 Sum::Right(_) => {
182 assert((ptr_addr & tag) < ptr_addr) by (bit_vector)
183 requires
184 ptr_addr & tag == tag,
185 (ptr_addr & !tag) != 0,
186 ;
187 },
188 }
189 }
190 let (is_right, real_ptr) = unsafe { remove_bits(ptr, 1 << Self::ALIGN_BITS) };
193
194 if is_right == 0 {
195 Either::Left(unsafe { L::from_raw(real_ptr.cast(), Tracked(perm.tracked_take_left())) })
198 } else {
199 Either::Right(
202 unsafe { R::from_raw(real_ptr.cast(), Tracked(perm.tracked_take_right())) },
203 )
204 }
205 }
206
207 open spec fn ptr_perm_match(ptr: *mut Self::Target, perm: Self::Permission) -> bool {
208 let tag = 1usize << Self::ALIGN_BITS;
209 match perm {
210 Sum::Left(left) => {
211 &&& ptr.addr() & tag == 0
212 &&& L::ptr_perm_match(ptr.cast(), left)
213 },
214 Sum::Right(right) => {
215 let untagged_ptr = ptr.with_addr((ptr.addr() & !tag));
216 let right_nonnull = nonnull_from_ptr_mut_spec(untagged_ptr);
217 &&& ptr.addr() & tag == tag
218 &&& (ptr.addr() & !tag) != 0
219 &&& R::ptr_perm_match(right_nonnull.cast().view_ptr_mut(), right)
220 },
221 }
222 }
223
224 open spec fn rel_perm(self, perm: Self::Permission) -> bool {
225 match (self, perm) {
226 (Either::Left(left), Sum::Left(left_perm)) => left.rel_perm(left_perm),
227 (Either::Right(right), Sum::Right(right_perm)) => right.rel_perm(right_perm),
228 _ => false,
229 }
230 }
231
232 axiom fn lemma_align_bits_range()
233 ensures
234 Self::ALIGN_BITS == if L::ALIGN_BITS < R::ALIGN_BITS {
235 L::ALIGN_BITS - 1
236 } else {
237 R::ALIGN_BITS - 1
238 },
239 ;
240}
241
242unsafe impl<'a, L: NonNullPtrRef<'a>, R: NonNullPtrRef<'a>> NonNullPtrRef<'a> for Either<L, R> {
243 type Ref = Either<L::Ref, R::Ref>;
244
245 type RefPermission = Sum<L::RefPermission, R::RefPermission>;
246
247 open spec fn ref_perm_view_permission(perm: Self::RefPermission) -> Self::Permission {
248 match perm {
249 Sum::Left(left) => Sum::Left(L::ref_perm_view_permission(left)),
250 Sum::Right(right) => Sum::Right(R::ref_perm_view_permission(right)),
251 }
252 }
253
254 open spec fn ref_rel_perm(r: Self::Ref, perm: Self::RefPermission) -> bool {
255 true
256 }
257
258 proof fn lemma_ref_perm_inv_impl_perm_inv(perm: Self::RefPermission) {
259 match perm {
260 Sum::Left(left) => L::lemma_ref_perm_inv_impl_perm_inv(left),
261 Sum::Right(right) => R::lemma_ref_perm_inv_impl_perm_inv(right),
262 }
263 }
264
265 proof fn borrow_ref_perm(tracked perm: &Self::RefPermission) -> (tracked ret:
266 Self::RefPermission) {
267 if perm is Left {
268 Sum::Left(L::borrow_ref_perm(perm.tracked_borrow_left()))
269 } else {
270 Sum::Right(R::borrow_ref_perm(perm.tracked_borrow_right()))
271 }
272 }
273
274 proof fn borrow_perm_as_ref_perm(tracked perm: &'a Self::Permission) -> (tracked ret:
275 Self::RefPermission) {
276 if perm is Left {
277 Sum::Left(L::borrow_perm_as_ref_perm(perm.tracked_borrow_left()))
278 } else {
279 Sum::Right(R::borrow_perm_as_ref_perm(perm.tracked_borrow_right()))
280 }
281 }
282
283 unsafe fn raw_as_ref(
284 raw: NonNull<Self::Target>,
285 Tracked(perm): Tracked<Self::RefPermission>,
286 ) -> Self::Ref {
287 proof_decl! {
288 let ghost align_bits = Self::ALIGN_BITS;
289 let ghost tag = 1usize << align_bits;
290 let ghost raw_addr = raw.view_ptr_mut()@.addr;
291 }
292 proof! {
293 Self::lemma_align_bits_range();
294 if perm is Left {
295 assert((raw_addr & !tag) == raw_addr) by (bit_vector)
296 requires
297 raw_addr & tag == 0,
298 ;
299 } else {
300 assert((raw_addr & tag) < raw_addr) by (bit_vector)
301 requires
302 raw_addr & tag == tag,
303 (raw_addr & !tag) != 0,
304 ;
305 }
306 }
307 let (is_right, real_ptr) = unsafe { remove_bits(raw, 1 << Self::ALIGN_BITS) };
310
311 if is_right == 0 {
312 proof!{
313 if perm is Right {
314 assert(tag != 0) by (bit_vector)
315 requires
316 tag == 1usize << align_bits,
317 align_bits < usize::BITS,
318 ;
319 assert(false);
320 }
321 }
322 Either::Left(
325 unsafe { L::raw_as_ref(real_ptr.cast(), Tracked(perm.tracked_take_left())) },
326 )
327 } else {
328 Either::Right(
331 unsafe { R::raw_as_ref(real_ptr.cast(), Tracked(perm.tracked_take_right())) },
332 )
333 }
334 }
335
336 #[verifier::spinoff_prover]
337 fn ref_as_raw(ptr_ref: Self::Ref) -> (NonNull<Self::Target>, Tracked<Self::RefPermission>) {
338 proof!{
339 Self::lemma_align_bits_range();
340 }
341 proof_decl!{
342 let ghost align_bits = Self::ALIGN_BITS;
343 let ghost tag = 1usize << align_bits;
344 let ghost l_align_bits = L::ALIGN_BITS;
345 let ghost r_align_bits = R::ALIGN_BITS;
346 }
347 proof!{
348 L::lemma_align_bits_range();
349 R::lemma_align_bits_range();
350 vstd::bits::lemma_usize_pow2_no_overflow(align_bits as nat);
351 vstd::bits::lemma_usize_pow2_no_overflow(l_align_bits as nat);
352 vstd::bits::lemma_usize_pow2_no_overflow(r_align_bits as nat);
353 vstd::bits::lemma_usize_shl_is_mul(1, align_bits as usize);
354 vstd::bits::lemma_usize_shl_is_mul(1, l_align_bits as usize);
355 vstd::bits::lemma_usize_shl_is_mul(1, r_align_bits as usize);
356 }
357 match ptr_ref {
358 Either::Left(left) => {
359 let (ptr, Tracked(perm)) = L::ref_as_raw(left);
361 proof! {
362 let ghost ptr_addr = ptr.view_ptr_mut().addr();
363 L::lemma_ref_perm_inv_impl_perm_inv(perm);
364 let extra_bits: u32 = (l_align_bits - align_bits) as u32;
365 let scale = 1usize << extra_bits;
366 vstd::bits::lemma_usize_pow2_no_overflow(extra_bits as nat);
367 vstd::bits::lemma_usize_shl_is_mul(1, extra_bits as usize);
368 vstd::arithmetic::power2::lemma_pow2_adds(align_bits as nat, extra_bits as nat);
369 assert(ptr_addr % tag == 0) by {
370 let big = 1usize << l_align_bits;
371 let q = ptr_addr / big;
372 vstd::arithmetic::div_mod::lemma_fundamental_div_mod(ptr_addr as int, big as int);
373 assert(ptr_addr == (q * scale) * tag) by (nonlinear_arith)
374 requires
375 ptr_addr == q * big,
376 big == tag * scale,
377 ;
378 vstd::arithmetic::div_mod::lemma_mod_multiples_basic(q * scale, tag as int);
379 };
380 assert(ptr_addr & tag == 0) by (bit_vector)
381 requires
382 ptr_addr % (1usize << l_align_bits) == 0,
383 tag == 1usize << align_bits,
384 align_bits < l_align_bits < usize::BITS,
385 ;
386 }
387 (ptr.cast(), Tracked(Sum::Left(perm)))
388 },
389 Either::Right(right) => {
390 let (ptr, Tracked(perm)) = R::ref_as_raw(right);
394 proof! {
395 Self::lemma_align_bits_range();
396 }
397 let tagged_ptr = ptr.map_addr_v(
398 |addr: NonZeroUsize| -> (ret: NonZeroUsize)
399 ensures
400 ret@ == addr@ | (1usize << Self::ALIGN_BITS),
401 {
402 proof {
403 let tag = 1usize << Self::ALIGN_BITS;
404 let a = addr@;
405 assert(a | tag != 0) by (bit_vector)
406 requires
407 a != 0,
408 ;
409 }
410 addr | 1usize << Self::ALIGN_BITS
411 },
412 );
413 proof! {
414 let ghost ptr_addr = ptr.view_ptr_mut().addr();
415 let ghost tagged_addr = tagged_ptr.view_ptr_mut().addr();
416 R::lemma_ref_perm_inv_impl_perm_inv(perm);
417 assert(tagged_addr & tag == tag) by (bit_vector)
418 requires
419 tagged_addr == ptr_addr | tag,
420 tag == 1usize << align_bits,
421 align_bits < r_align_bits < usize::BITS,
422 ptr_addr % (1usize << r_align_bits) == 0,
423 ptr_addr != 0,
424 ;
425 assert(ptr_addr & tag == 0) by (bit_vector)
426 requires
427 ptr_addr % (1usize << r_align_bits) == 0,
428 tag == 1usize << align_bits,
429 align_bits < r_align_bits < usize::BITS,
430 ;
431 assert(tagged_addr & !tag == ptr_addr) by (bit_vector)
432 requires
433 tagged_addr == ptr_addr | tag,
434 ptr_addr & tag == 0,
435 ;
436 let extra_bits: u32 = (r_align_bits - align_bits) as u32;
437 let scale = 1usize << extra_bits;
438 vstd::bits::lemma_usize_pow2_no_overflow(extra_bits as nat);
439 vstd::bits::lemma_usize_shl_is_mul(1, extra_bits as usize);
440 vstd::arithmetic::power2::lemma_pow2_adds(align_bits as nat, extra_bits as nat);
441 assert(tagged_addr == ptr_addr + tag) by (bit_vector)
442 requires
443 tagged_addr == ptr_addr | tag,
444 ptr_addr & tag == 0,
445 ;
446 assert(ptr_addr % tag == 0) by {
447 let big = 1usize << r_align_bits;
448 let q = ptr_addr / big;
449 vstd::arithmetic::div_mod::lemma_fundamental_div_mod(ptr_addr as int, big as int);
450 assert(ptr_addr == (q * scale) * tag) by (nonlinear_arith)
451 requires
452 ptr_addr == q * big,
453 big == tag * scale,
454 ;
455 vstd::arithmetic::div_mod::lemma_mod_multiples_basic(q * scale, tag as int);
456 };
457 assert(tagged_addr % tag == 0) by {
458 vstd::arithmetic::div_mod::lemma_mod_add_multiples_vanish(ptr_addr as int, tag as int);
459 }
460 }
461 (tagged_ptr.cast(), Tracked(Sum::Right(perm)))
462 },
463 }
464 }
465}
466
467} #[verus_verify(dual_spec)]
470const fn min(a: u32, b: u32) -> u32 {
471 if a < b { a } else { b }
472}
473
474verus! {
475
476#[verus_spec(ret =>
482 requires
483 (ptr.view_ptr_mut().addr() & bits) < ptr.view_ptr_mut().addr(),
484 (ptr.view_ptr_mut().addr() & !bits) != 0,
485 ensures
486 ret.0 == (ptr.view_ptr_mut().addr() & bits),
487 ret.1.view_ptr_mut() == ptr.view_ptr_mut().with_addr((ptr.view_ptr_mut().addr() & !bits) as usize),
488)]
489unsafe fn remove_bits<T>(ptr: NonNull<T>, bits: usize) -> (usize, NonNull<T>) {
490 use core::num::NonZeroUsize;
491
492 let removed_bits = ptr.addr_v().get() & bits;
493 let result_ptr = ptr.map_addr_v(
494 |addr| -> (ret: NonZeroUsize)
495 requires
496 addr@ & !bits != 0,
497 ensures
498 ret@ == addr@ & !bits,
499 {
500 unsafe { NonZeroUsize::new_unchecked(addr.get() & !bits) }
502 },
503 );
504 (removed_bits, result_ptr)
505}
506
507#[verifier::spinoff_prover]
508proof fn lemma_aligned_addr_clears_tag_bit(
509 addr: usize,
510 tag: usize,
511 align_bits: u32,
512 ptr_align_bits: u32,
513)
514 requires
515 addr % (1usize << ptr_align_bits) == 0,
516 tag == 1usize << align_bits,
517 align_bits < ptr_align_bits < usize::BITS,
518 ensures
519 addr & tag == 0,
520{
521 assert(addr & (1usize << align_bits) == 0) by (bit_vector)
522 requires
523 addr % (1usize << ptr_align_bits) == 0,
524 0u32 < ptr_align_bits,
525 align_bits < ptr_align_bits,
526 ptr_align_bits <= 63u32,
527 ;
528 assert(addr & tag == 0) by (bit_vector)
529 requires
530 addr & (1usize << align_bits) == 0,
531 tag == 1usize << align_bits,
532 align_bits < 64u32,
533 ;
534}
535
536} #[cfg(ktest)]
538mod test {
539 use alloc::{boxed::Box, sync::Arc};
540
541 use super::*;
542 use crate::{prelude::ktest, sync::RcuOption};
543
544 type Either32 = Either<Arc<u32>, Box<u32>>;
545 type Either16 = Either<Arc<u32>, Box<u16>>;
546
547 #[ktest]
548 fn alignment() {
549 assert_eq!(<Either32 as NonNullPtr>::ALIGN_BITS, 1);
550 assert_eq!(<Either16 as NonNullPtr>::ALIGN_BITS, 0);
551 }
552
553 #[ktest]
554 fn left_pointer() {
555 let val: Either16 = Either::Left(Arc::new(123));
556
557 let ptr = NonNullPtr::into_raw(val);
558 assert_eq!(ptr.addr().get() & 1, 0);
559
560 let ref_ = unsafe { <Either16 as NonNullPtr>::raw_as_ref(ptr) };
561 assert!(matches!(ref_, Either::Left(ref r) if ***r == 123));
562
563 let ptr2 = <Either16 as NonNullPtr>::ref_as_raw(ref_);
564 assert_eq!(ptr, ptr2);
565
566 let val = unsafe { <Either16 as NonNullPtr>::from_raw(ptr) };
567 assert!(matches!(val, Either::Left(ref r) if **r == 123));
568 drop(val);
569 }
570
571 #[ktest]
572 fn right_pointer() {
573 let val: Either16 = Either::Right(Box::new(456));
574
575 let ptr = NonNullPtr::into_raw(val);
576 assert_eq!(ptr.addr().get() & 1, 1);
577
578 let ref_ = unsafe { <Either16 as NonNullPtr>::raw_as_ref(ptr) };
579 assert!(matches!(ref_, Either::Right(ref r) if ***r == 456));
580
581 let ptr2 = <Either16 as NonNullPtr>::ref_as_raw(ref_);
582 assert_eq!(ptr, ptr2);
583
584 let val = unsafe { <Either16 as NonNullPtr>::from_raw(ptr) };
585 assert!(matches!(val, Either::Right(ref r) if **r == 456));
586 drop(val);
587 }
588
589 #[ktest]
590 fn rcu_store_load() {
591 let rcu: RcuOption<Either32> = RcuOption::new_none();
592 assert!(rcu.read().get().is_none());
593
594 rcu.update(Some(Either::Left(Arc::new(888))));
595 assert!(matches!(rcu.read().get().unwrap(), Either::Left(r) if **r == 888));
596
597 rcu.update(Some(Either::Right(Box::new(999))));
598 assert!(matches!(rcu.read().get().unwrap(), Either::Right(r) if **r == 999));
599 }
600}