predicate sorted(xs: seq<int>)
{
  forall i :: 0 <= i && i + 1 < |xs| ==> xs[i] <= xs[i + 1]
}


method Partition(xs: seq<int>, pivot: int) returns (lower: seq<int>, upper: seq<int>)
  ensures multiset(lower) + multiset(upper) == multiset(xs)
  ensures  forall x :: x in lower ==> x <= pivot
  ensures  forall x :: x in upper ==> x > pivot
  ensures |lower| + |upper| == |xs| // how does this help?
{
  if |xs| == 0 {
    lower, upper := [], [];
  } else {
    var lowerRest, upperRest := Partition(xs[1..], pivot);

    if xs[0] <= pivot {
      lower, upper := [xs[0]] + lowerRest, upperRest;
    } else {
      lower, upper := lowerRest, [xs[0]] + upperRest;
    }

    assert xs == [xs[0]] + xs[1..]; // What does this assertion do?
  }
}



method QuickSort(xs: seq<int>) returns (result: seq<int>)
  ensures sorted(result)
  ensures multiset(result) == multiset(xs)
  decreases |xs|
{
  if |xs| == 0 {
    result := [];
  } else {
    var pivot := xs[0];
    var lower, upper := Partition(xs[1..], pivot);

	assert(xs == [xs[0]] + xs[1..]); // why is this necessary?

    var sortedLower := QuickSort(lower);
    var sortedUpper := QuickSort(upper);
	
	Preservation(lower, sortedLower, upper, sortedUpper);
    
	result := sortedLower + [pivot] + sortedUpper;
	
	Preservation(lower, sortedLower, [pivot], [pivot]);
	SortedConcatenation(sortedLower, [pivot]);
	Preservation([pivot], [pivot], upper, sortedUpper);
	SortedConcatenation(sortedLower + [pivot], sortedUpper);
  }
}


lemma Preservation(s: seq<int>, s': seq<int>, t: seq<int>, t': seq<int>)
	requires forall x,y :: (x in s && y in t) ==> x <= y
	requires multiset(s) == multiset(s')
	requires multiset(t) == multiset(t')
	ensures forall x,y :: (x in s' && y in t') ==> x <= y
{
	//proof by contradiction
	if !(forall x,y :: (x in s' && y in t') ==> x <= y) {
		var x,y :| x in s' && y in t' && x > y;
		calc ==> {
			x in s';
			x in multiset(s');
			x in multiset(s);
			x in s;
		}
		calc ==> {
			y in t';
			y in multiset(t');
			y in multiset(t);
			y in t;
		}
	}
}


lemma SortedConcatenation(s: seq<int>, t: seq<int>) 
	requires sorted(s) && sorted(t)
	requires forall x,y :: (x in s && y in t) ==> x <= y
	ensures sorted(s + t)
{
	// proof by contradiction
	if !sorted(s + t) {
		assert(exists i :: 0 <= i && i < |s+t| - 1 && (s+t)[i] > (s+t)[i+1]);
		var i :| 0 <= i && i < |s+t| - 1 && (s+t)[i] > (s+t)[i+1];
		if i < |s| - 1 {

		} else if i > |s| - 1{

		} else {
			assert((s+t)[i] in s);
			assert((s+t)[i+1] in t);
		}
	}
}