use vstd::prelude::*;
verus! {
pub struct Sampler {
pub num_items: usize,
pub sample_size: usize,
pub distribution: Vec<u64>,
pub selected: Vec<usize>,
}
impl Sampler {
pub open spec fn all_valid(s: Seq<usize>, num_items: usize) -> bool {
forall|i: int| 0 <= i < s.len() ==> #[trigger] s[i] < num_items
}
pub open spec fn all_distinct(s: Seq<usize>) -> bool {
forall|i: int, j: int|
0 <= i < s.len() && 0 <= j < s.len() && i != j ==> s[i] != s[j]
}
pub open spec fn contains(&self, n: usize) -> bool {
exists|i: int| 0 <= i < self.selected.len() && self.selected@[i] == n
}
pub open spec fn type_invariant(&self) -> bool {
&&& self.distribution.len() == self.num_items
&&& Self::all_valid(self.selected@, self.num_items)
&&& Self::all_distinct(self.selected@)
}
pub open spec fn bounded_sample(&self) -> bool {
self.selected.len() <= self.sample_size
}
pub open spec fn support_consistency(&self) -> bool {
forall|i: int|
0 <= i < self.selected.len()
==> #[trigger] self.distribution@[self.selected@[i] as int] > 0
}
pub open spec fn inv(&self) -> bool {
self.type_invariant() && self.bounded_sample() && self.support_consistency()
}
pub fn new(distribution: Vec<u64>, sample_size: usize) -> (s: Sampler)
ensures
s.num_items == distribution@.len(),
s.sample_size == sample_size,
s.distribution@ == distribution@,
s.selected@.len() == 0,
s.inv(),
{
let num_items = distribution.len();
Sampler { num_items, sample_size, distribution, selected: Vec::new() }
}
pub fn contains_exec(&self, n: usize) -> (b: bool)
ensures b == self.contains(n),
{
let len = self.selected.len();
let mut i: usize = 0;
while i < len
invariant
i <= len,
len == self.selected.len(),
forall|k: int| 0 <= k < i ==> self.selected@[k] != n,
decreases len - i,
{
if self.selected[i] == n {
assert(self.selected@[i as int] == n);
return true;
}
i = i + 1;
}
assert(!self.contains(n));
false
}
pub fn sample(&mut self, i: usize)
requires
old(self).inv(),
old(self).selected.len() < old(self).sample_size, i < old(self).num_items, old(self).distribution@[i as int] > 0, !old(self).contains(i), ensures
final(self).num_items == old(self).num_items,
final(self).sample_size == old(self).sample_size,
final(self).distribution@ == old(self).distribution@,
final(self).selected@ == old(self).selected@.push(i),
final(self).inv(),
{
let ghost os = self.selected@;
self.selected.push(i);
assert(self.selected@ == os.push(i));
assert(Self::all_valid(self.selected@, self.num_items));
assert(Self::all_distinct(self.selected@)) by {
assert forall|a: int, b: int|
0 <= a < self.selected@.len() && 0 <= b < self.selected@.len() && a != b
implies self.selected@[a] != self.selected@[b] by {
if a < os.len() && b < os.len() {
} else if a == os.len() && b < os.len() {
assert(self.selected@[b] == os[b]);
assert(os[b] != i); } else if b == os.len() && a < os.len() {
assert(self.selected@[a] == os[a]);
assert(os[a] != i);
}
}
};
assert(self.support_consistency()) by {
assert forall|a: int| 0 <= a < self.selected@.len()
implies #[trigger] self.distribution@[self.selected@[a] as int] > 0 by {
if a < os.len() {
assert(self.selected@[a] == os[a]); } else {
assert(self.selected@[a] == i); }
}
};
}
pub fn zero(&mut self, i: usize) -> (ok: bool)
requires old(self).inv(),
ensures
final(self).inv(),
final(self).num_items == old(self).num_items,
final(self).sample_size == old(self).sample_size,
ok == (i < old(self).num_items && !old(self).contains(i)),
ok ==> final(self).distribution@ == old(self).distribution@.update(i as int, 0),
!ok ==> final(self).distribution@ == old(self).distribution@,
final(self).selected@ == old(self).selected@,
{
if i >= self.num_items || self.contains_exec(i) {
false
} else {
self.distribution.set(i, 0);
true
}
}
pub fn draw_weighted(&mut self, i: usize, r: u64) -> (accepted: bool)
requires
old(self).inv(),
i < old(self).num_items,
ensures
final(self).inv(),
final(self).num_items == old(self).num_items,
final(self).sample_size == old(self).sample_size,
final(self).distribution@ == old(self).distribution@,
accepted ==> final(self).selected@ == old(self).selected@.push(i),
!accepted ==> final(self).selected@ == old(self).selected@,
{
if self.selected.len() >= self.sample_size {
return false; }
if r >= self.distribution[i] {
return false; }
if self.contains_exec(i) {
return false; }
self.sample(i);
true
}
pub fn draw_uniform(&mut self, i: usize) -> (accepted: bool)
requires
old(self).inv(),
i < old(self).num_items,
ensures
final(self).inv(),
final(self).num_items == old(self).num_items,
final(self).sample_size == old(self).sample_size,
final(self).distribution@ == old(self).distribution@,
accepted ==> final(self).selected@ == old(self).selected@.push(i),
!accepted ==> final(self).selected@ == old(self).selected@,
{
if self.selected.len() >= self.sample_size {
return false;
}
if self.distribution[i] == 0 {
return false; }
if self.contains_exec(i) {
return false;
}
self.sample(i);
true
}
}
}