6.S057: 6.S057: Verified Software Engineering

Class 7

class07.dfy
dfy
// We turn our attention to some functional algorithms.
// Our next two algorithms will be sorting algorithms. So let's start
// by considering the specification of sorting.

// We'll be working on our familiar lists.
datatype List<T> = Nil | Cons(head: T, tail: List<T>)

function Length<T>(xs: List<T>): nat {
  match xs
  case Nil => 0
  case Cons(_, tail) => 1 + Length(tail)
}

function Append<X>(xs: List<X>, ys: List<X>): List<X>
  ensures Length(Append(xs, ys)) == Length(xs) + Length(ys)
{
  match xs
  case Nil => ys
  case Cons(x, tail) => Cons(x, Append(tail, ys))
}

predicate Member<X(==)>(x: X, xs: List<X>) {
  match xs
  case Nil => false
  case Cons(y, xs') => x == y || Member(x, xs')
}

function At<X>(xs: List<X>, i: nat): X
  requires i < Length(xs)
{
  if i == 0 then xs.head else At(xs.tail, i - 1)
}

// In the general case, we want to define sorting generically, that is, for any type
// of list. We can do this by using an `lt` function as we did in the previous class.
// However, to focus on the algorithms rather than properties of `lt`, we'll restrict
// consideration of any ordering matters to integers for now.
type IntList = List<int>

// To determine whether or not a list is sorted, we need to look at more
// than one element at a time. Here is one way we can do that.
predicate Ordered(xs: IntList) {
  match xs
  case Nil => true
  case Cons(_, Nil) => true
  case Cons(x, Cons(y, _)) => x <= y && Ordered(xs.tail)
}

// Remark: Since match-cases are ordered, a shorter way to write the same function is
predicate Ordered'(xs: IntList) {
  match xs
  case Cons(x, Cons(y, _)) => x <= y && Ordered'(xs.tail)
  case _ => true
}

// When we write a specification, we can make mistakes, just like we can make mistakes
// when writing a program. One way to test a specification is to prove some properties
// about it, to make sure it satisfies some properties we'd expect. Here are two such
// properties:

// The smallest element of a list is at its head
lemma HeadIsSmallest(xs: IntList, y: int)
  requires Ordered(xs) && xs.Cons?
  requires Member(y, xs)
  ensures xs.head <= y
{
}

// More generally, we can say that any two elements in a list are ordered according to
// their positions.
lemma AllOrdered(xs: IntList, i: nat, j: nat)
  requires Ordered(xs) && i <= j < Length(xs)
  ensures At(xs, i) <= At(xs, j)
{
  if i != 0 {
    // let's move closer to the elements i and j
    AllOrdered(xs.tail, i - 1, j - 1);
  } else if i == j {
    // easy
  } else {
    AllOrdered(xs.tail, 0, j - 1);
  }
}

// Just because a function returns an ordered list does not make it a sorting function.
// A sorting function is expected to return the same elements as it's given, just possibly
// in a different order. There are several ways to specify this "same elements" property.
// One way is to count the number of occurrences of each possible element and to say
// that these counts agree for the input and output lists, for every possible value.
function Count<X(==)>(xs: List<X>, val: X): nat {
  match xs
  case Nil => 0
  case Cons(x, tail) => (if val == x then 1 else 0) + Count(tail, val)
}

// Using Ordered and Count, we can describe what it means for a function to be a sorting
// function. Let's write our first sorting function using the "insertion sort" algorithm.
function InsertionSort(xs: IntList): IntList {
  match xs
  case Nil => Nil
  case Cons(x, tail) =>
    var sortedTail := InsertionSort(tail);
    Insert(x, sortedTail)
}

function Insert(x: int, xs: IntList): IntList {
  match xs
  case Nil => Cons(x, Nil)
  case Cons(y, tail) =>
    if x <= y then Cons(x, xs) else Cons(y, Insert(x, tail))
}
// Note that Insert is defined for any IntList, not just a sorted one. However, we will
// only call it on a sorted list, so when it comes time to proving something about Insert,
// we will restrict the lemma to sorted lists.

// To prove the correctness of InsertionSort, we'll have to show that it outputs
// an ordered list (using `Ordered`) and that it preserves element counts (using
// `Count`).

// Here is the lemma for the first of those two tasks:
lemma InsertionSortOrdered(xs: IntList)
  ensures Ordered(InsertionSort(xs))
{
  match xs
  case Nil =>
    // easy
  case Cons(x, tail) =>
    // When functions get complicated, it can give us peace of mind to start by
    // writing down what the function we're trying to prove something about does.
    var sortedTail := InsertionSort(tail);
    var result := Insert(x, sortedTail);
    assert result == InsertionSort(xs); // this assertion makes sure we copied the
                                        //     previous two lines correctly
    // We can start writing out the proof in detail. But if we think the situation
    // looks pretty simple, we can start by calling the induction hypothesis in a way
    // analogous to how the function is used.
    InsertionSortOrdered(tail);

    // This was not enough. Not too surprising, actually. We'll need to know something
    // about what Insert does. We'll define a lemma InsertOrdered (below) and call it
    // from here.
    InsertOrdered(x, sortedTail);

    // That did the trick!
}

lemma InsertOrdered(x: int, xs: IntList)
  requires Ordered(xs)
  ensures Ordered(Insert(x, xs))
{
  // This proof is completed by automatic induction.
}

// Here is a lemma for the second task, showing that InsertionSort preserves element
// counts. The easiest way to write such a lemma is to parameterize it by an arbitrary
// element.
lemma InsertionSortPreservesCounts(xs: IntList, val: int)
  ensures Count(xs, val) == Count(InsertionSort(xs), val)
{
  match xs
  case Nil =>
  case Cons(x, tail) =>
    var sortedTail := InsertionSort(tail);
    var result := Insert(x, sortedTail);
    assert result == InsertionSort(xs);

    calc {
      Count(InsertionSort(xs), val);
      Count(Insert(x, sortedTail), val);
      { InsertCount(x, sortedTail, val); }
      (if val == x then 1 else 0) + Count(sortedTail, val);
      { InsertionSortPreservesCounts(tail, val); }
      (if val == x then 1 else 0) + Count(tail, val);
      Count(xs, val);
    }
}

lemma InsertCount(x: int, xs: IntList, val: int)
  ensures Count(Insert(x, xs), val) == (if val == x then 1 else 0) + Count(xs, val)
{
}

// That concludes the proof of correctness of InsertionSort. InsertionSort has an
// O(n^2) running-time complexity. Next, let's write and prove MergeSort, which has
// an O(n log n) running-time complexity.

// MergeSort proceeds by splitting the input into two pieces, sorting each one, and
// then merging the two sorted lists. We start by writing a function for splitting
// a list into two pieces.
function Split'<X>(xs: List<X>, n: nat): (List<X>, List<X>)
  requires n <= Length(xs)
{
  if n == 0 then
    (Nil, xs)
  else
    var (a, b) := Split'(xs.tail, n - 1); // To access `xs.tail`, we must know `xs.Cons?`.
    (Cons(xs.head, a), b)                 //     How do we know `xs.Cons?` holds here?    
}

// Here is a lemma that states the correctness of Split.
lemma AboutSplit'<X>(xs: List<X>, n: nat)
  requires n <= Length(xs)
  ensures var (a, b) := Split'(xs, n);
    Append(a, b) == xs && Length(a) == n
{
}

// Our every call to Split will need the properties stated by AboutSplit'.
// Therefore, analogously to what we did for Append, we'll use a function postcondition
// to build these properties into the definition of Split.
function Split<X>(xs: List<X>, n: nat): (List<X>, List<X>)
  requires n <= Length(xs)
  ensures var (a, b) := Split(xs, n);
    Append(a, b) == xs && Length(a) == n
{
  if n == 0 then
    (Nil, xs)
  else
    var (a, b) := Split(xs.tail, n - 1);
    (Cons(xs.head, a), b)
}

// Now, we can write MergeSort. It needs to have the length of the given list,
// so that it can call Split with half of that length. MergeSort could compute
// this length with every call. But each recursive call already knows the length
// of the list it passes in, so we'll define an auxiliary MergeSort function
// that takes the length as a parameter and does most of the work.
function MergeSort(xs: IntList): IntList {
  MergeSortAux(Length(xs), xs)
}

function MergeSortAux(len: nat, xs: IntList): IntList
  requires len == Length(xs)
{
  if len < 2 then
    xs // such a short list is already sorted
  else
    var (left, right) := Split(xs, len / 2);
    var leftSorted := MergeSortAux(len / 2, left);
    var rightSorted := MergeSortAux(len - len / 2, right);
    Merge(leftSorted, rightSorted)
}

// We need to write Merge.
function Merge(xs: IntList, ys: IntList): IntList {
  // Here, it's convenient to do the pattern match on the pair (xs, ys).
  match (xs, ys)
  case (Nil, _) => ys
  case (_, Nil) => xs
  case (Cons(x, xs'), Cons(y, ys')) =>
    if x <= y then
      Cons(x, Merge(xs', ys))
    else
      Cons(y, Merge(xs, ys'))
}
// Note that, like the auxiliary function Insert that we saw with InsertionSort,
// function Merge is defined for any lists, not just sorted lists. But since we
// only ever call Merge on sorted lists, we will feel free to restrict lemmas
// about Merge to cases where the two given lists are sorted.

// We'll now prove the two properties about MergeSort (or, rather, of MergeSortAux),
// that it returns an ordered list and that it preserved the elements of the list.
// For each of these properties, we'll need a corresponding auxiliary lemma about
// the auxiliary function Merge.
//
// To proceed in an orderly fashion, we start by writing the structure of the lemma,
// following the structure of function MergeSortAux. This makes things more verbose
// for now, but that can be cleaned up once we're done with the proof.
// Hint: While writing down this structure, it can be convenient to temporarily
// comment out the postcondition. This will make our assertions the only things
// the verifier has to deal with. 
lemma MergeSortOrdered'(len: nat, xs: IntList)
  requires len == Length(xs)
//  ensures Ordered(MergeSort(xs))
{
  var mergeSort := MergeSortAux(len, xs);
  if len < 2 {
    assert mergeSort == xs;
  } else {
    var (left, right) := Split(xs, len / 2);
    var leftSorted := MergeSortAux(len / 2, left);
    var rightSorted := MergeSortAux(len - len / 2, right);
    assert mergeSort == Merge(leftSorted, rightSorted);
  }
}

lemma MergeSortOrdered(len: nat, xs: IntList)
  requires len == Length(xs)
  ensures Ordered(MergeSort(xs))
{
  var mergeSort := MergeSortAux(len, xs);
  if len < 2 {
    assert mergeSort == xs;
  } else {
    var (left, right) := Split(xs, len / 2);
    var leftSorted := MergeSortAux(len / 2, left);
    var rightSorted := MergeSortAux(len - len / 2, right);
    assert mergeSort == Merge(leftSorted, rightSorted);

    // We expect we'll need the Ordered property for the two
    // recursive calls, so we'll call the lemma recursively for
    // those.
    MergeSortOrdered(len / 2, left);
    MergeSortOrdered(len - len / 2, right);

    // We'll also need a lemma about Merge returning something ordered (below).
    MergeOrdered(leftSorted, rightSorted);
  }
}

lemma MergeOrdered(xs: IntList, ys: IntList)
  requires Ordered(xs) && Ordered(ys)
  ensures Ordered(Merge(xs, ys))
{
}

// Next, we prove that MergeSort preserves the elements of the input list.
lemma MergeSortPreservesCounts(len: nat, xs: IntList, val: int)
  requires len == Length(xs)
  ensures Count(xs, val) == Count(MergeSortAux(len, xs), val)
{
  var mergeSort := MergeSortAux(len, xs);
  if len < 2 {
    assert mergeSort == xs;
  } else {
    var (left, right) := Split(xs, len / 2);
    var leftSorted := MergeSortAux(len / 2, left);
    var rightSorted := MergeSortAux(len - len / 2, right);
    assert mergeSort == Merge(leftSorted, rightSorted);

    // There are several steps and operations involved, so let's take it
    // one step at a time using a "calc" statement.
    calc {
      Count(MergeSortAux(len, xs), val);
      Count(Merge(leftSorted, rightSorted), val);
      { MergePreservesCounts(leftSorted, rightSorted, val); }
      Count(leftSorted, val) + Count(rightSorted, val);
      { MergeSortPreservesCounts(len / 2, left, val); }
      Count(left, val) + Count(rightSorted, val);
      { MergeSortPreservesCounts(len - len / 2, right, val); }
      Count(left, val) + Count(right, val);
      // Here, we notice we need to be able to get from these two counts to
      // a count of where "left" and "right" came from, before the Split.
      // We know from the postcondition of Split that the two pieces returned,
      // when concatenated, equal the larger part that was split. So, what
      // we need is a lemma that relates Count and Append. We'll add that
      // lemma below and use it here.
      { CountAppend(left, right, val); }
      Count(Append(left, right), val);
      Count(xs, val);
    }
  }
}

lemma MergePreservesCounts(xs: IntList, ys: IntList, val: int)
  ensures Count(Merge(xs, ys), val) == Count(xs, val) + Count(ys, val)
{
}

lemma CountAppend<X>(xs: List<X>, ys: List<X>, val: X)
  ensures Count(Append(xs, ys), val) == Count(xs, val) + Count(ys, val)
{
}

Exercises

class07/list.dfy
dfy
datatype List<T> = Nil | Cons(head: T, tail: List<T>)

// Implement a useful list membership test...
predicate Member<X(==)>(x: X, xs: List<X>)
{
  match xs
  
}

// ... and implement a useful indexing function
function At<X>(xs: List<X>, i: nat): X
{
  xs.head
}
class07/ordered.dfy
dfy
datatype List<T> = Nil | Cons(head: T, tail: List<T>)

function Length<T>(xs: List<T>): nat {
  match xs
  case Nil => 0
  case Cons(_, tail) => 1 + Length(tail)
}

type IntList = List<int>

function At<X>(xs: List<X>, i: nat): X

predicate Ordered(xs: IntList)

lemma AllOrdered(xs: IntList, ...)
  requires ...
  ensures ...
{
  ...
}
class07/insertionsort.dfy
dfy
datatype List<T> = Nil | Cons(head: T, tail: List<T>)

type IntList = List<int>

function Count<X(==)>(xs: List<X>, val: X): nat {
  match xs
  case Nil => 0
  case Cons(x, tail) => (if val == x then 1 else 0) + Count(tail, val)
}

function InsertionSort(xs: IntList): IntList {
  match xs
  case Nil => Nil
  case Cons(x, tail) =>
    var sortedTail := InsertionSort(tail);
    Insert(x, sortedTail)
}

function Insert(x: int, xs: IntList): IntList {
  match xs
  case Nil => Cons(x, Nil)
  case Cons(y, tail) =>
    if x <= y then Cons(x, xs) else Cons(y, Insert(x, tail))
}

lemma InsertionSortPreservesCounts(xs: IntList, val: int)
  ensures Count(xs, val) == Count(InsertionSort(xs), val)
{
  match xs
  case Nil =>
  case Cons(x, tail) =>
    var sortedTail := InsertionSort(tail);
    var result := Insert(x, sortedTail);
    assert result == InsertionSort(xs);

    // Complete this case of the proof...

}

// ... perhaps with the help of a lemma about Insert and Count
class07/mergesort.dfy
dfy
datatype List<T> = Nil | Cons(head: T, tail: List<T>)

function Length<T>(xs: List<T>): nat {
  match xs
  case Nil => 0
  case Cons(_, tail) => 1 + Length(tail)
}

function Append<X>(xs: List<X>, ys: List<X>): List<X>
  ensures Length(Append(xs, ys)) == Length(xs) + Length(ys)
{
  match xs
  case Nil => ys
  case Cons(x, tail) => Cons(x, Append(tail, ys))
}

type IntList = List<int>

function Count<X(==)>(xs: List<X>, val: X): nat {
  match xs
  case Nil => 0
  case Cons(x, tail) => (if val == x then 1 else 0) + Count(tail, val)
}

function Split<X>(xs: List<X>, n: nat): (List<X>, List<X>)
  requires n <= Length(xs)
  ensures var (a, b) := Split(xs, n); Append(a, b) == xs && Length(a) == n
{
  if n == 0 then
    (Nil, xs)
  else
    var (a, b) := Split(xs.tail, n-1);
    (Cons(xs.head, a), b)
}

function MergeSort(xs: IntList): IntList {
  MergeSortAux(Length(xs), xs)
}

function MergeSortAux(len: nat, xs: IntList): IntList
  requires len == Length(xs)
{
  if len < 2 then
    xs // such a short list is already sorted
  else
    var (left, right) := Split(xs, len / 2);
    var leftSorted := MergeSortAux(len / 2, left);
    var rightSorted := MergeSortAux(len - len / 2, right);
    Merge(leftSorted, rightSorted)
}

function Merge(xs: IntList, ys: IntList): IntList {
  match (xs, ys)
  case (Nil, _) => ys
  case (_, Nil) => xs
  case (Cons(x, xs'), Cons(y, ys')) =>
    if x <= y then
      Cons(x, Merge(xs', ys))
    else
      Cons(y, Merge(xs, ys'))
}

lemma {:induction false} MergeSortPreservesCounts(len: nat, xs: IntList, val: int)
  requires len == Length(xs)
  ensures Count(xs, val) == Count(MergeSortAux(len, xs), val)
{
  var mergeSort := MergeSortAux(len, xs);
  if len < 2 {
    assert mergeSort == xs;
  } else {
    var (left, right) := Split(xs, len / 2);
    var leftSorted := MergeSortAux(len / 2, left);
    var rightSorted := MergeSortAux(len - len / 2, right);
    assert mergeSort == Merge(leftSorted, rightSorted);

    // Complete this case of the proof:
    // - notice that automatic induction is off
    // - feel free to add or remove intermediate steps
    calc {
      Count(MergeSortAux(len, xs), val);
      Count(Merge(leftSorted, rightSorted), val);
      { }
      Count(leftSorted, val) + Count(rightSorted, val);
      { }
      Count(left, val) + Count(rightSorted, val);
      { }
      Count(left, val) + Count(right, val);
      { }
      Count(xs, val);
    }
  }
}
Copyright 6.S057 course staff.