// 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 = Nil | Cons(head: T, tail: List) function Length(xs: List): nat { match xs case Nil => 0 case Cons(_, tail) => 1 + Length(tail) } function Append(xs: List, ys: List): List 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, xs: List) { match xs case Nil => false case Cons(y, xs') => x == y || Member(x, xs') } function At(xs: List, 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 // 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(xs: List, 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'(xs: List, n: nat): (List, List) 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'(xs: List, 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(xs: List, n: nat): (List, List) 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(xs: List, ys: List, val: X) ensures Count(Append(xs, ys), val) == Count(xs, val) + Count(ys, val) { }