/* XXX possibly buggy!  an example for testing.  can you find all the bugs? */

/* AVL Trees
 * 15-122 Principles of Imperative Computation, Spring 2011
 * Frank Pfenning
 */

#use <string>

/* Client-side implementation */

/* elements */
struct wcount {
  string word;			/* key */
  int count;			/* data = wordcount */
};
typedef struct wcount* wcount;

/* key comparison */
int compare(string s1, string s2) {
  return string_compare(s1,s2);
}

/* extracting keys from elements */
string elem_key(wcount wc)
//@requires wc != NULL;
{
  return wc->word;
}

/* Client-side interface declarations */
typedef string key;
typedef struct wcount* elem;	/* NULL must be an elem */
key elem_key(elem e);
int compare(key k1, key k2);

/* Library-side interface declarations */
typedef struct bst* bst;
bst bst_new();
void bst_insert(bst B, elem e)  /* replace if elem with same key as x in B */
//@requires e != NULL;
  ;
elem bst_search(bst B, key k);  /* return NULL if not in tree */

/* Implementation */
struct tree {
  int height;
  struct tree* left;
  elem data;
  struct tree* right;
};
typedef struct tree* tree;

struct bst {
  tree root;
};

/* is_ordered(T, lower, upper) checks if all elements in T
 * are strictly in the interval (elem_key(lower),elem_key(kupper)).
 * lower = NULL represents -infinity; upper = NULL represents +infinity
 */
bool is_ordered(tree T, elem lower, elem upper) {
  if (T == NULL) return true;
  if (T->data == NULL) return false;
  key k = elem_key(T->data);
  if (!(lower == NULL || compare(elem_key(lower),k) < 0))
    return false;
  if (!(upper == NULL || compare(k,elem_key(upper)) < 0))
    return false;
  return is_ordered(T->left, lower, T->data)
    && is_ordered(T->right, T->data, upper);
}

bool is_ordtree(tree T) {
  /* initially, we have no bounds - pass in NULL */
  return is_ordered(T, NULL, NULL);
}

int height(tree T)
{
  if (T == NULL)
    return 0;
  else
    return T->height;
}

tree make(tree l, elem e, tree r)
//@requires is_ordered(l, NULL, e);
//@requires is_ordered(r, e, NULL);
//@ensures is_ordtree(\result);
{
  int hl = height(l);
  int hr = height(r);
  int h = hl > hr ? hl+1 : hr+1;
  tree T = alloc(struct tree);
  T->height = h;
  T->left = l;
  T->data = e;
  T->right = r;
  return T;
}

bool is_balanced(tree T)
{
  if (T == NULL) return true;
  int h = T->height;
  int hl = height(T->left);
  int hr = height(T->right);
  if (!(h == (hl > hr ? hl+1 : hr+1))) return false;
  if (hl > hr+1 || hr > hl+1) return false;
  return is_balanced(T->left) && is_balanced(T->right);
}

int abs(int x)
//@requires x > (1<<31);
//@ensures \result >= 0;
{ return x < 0 ? -x : x; }

tree balance(tree l, elem e, tree r)
//@requires is_ordered(l, NULL, e);
//@requires is_ordered(r, e, NULL);
//@requires is_balanced(l) && is_balanced(r);
//@requires abs(height(l)-height(r)) <= 2;
//@ensures is_ordtree(\result);
//@ensures is_balanced(\result);
{
  int hl = height(l);
  int hr = height(r);
  if (hl > hr+1) {
    //@assert hl == hr+2;
    if (height(l->left) >= height(l->right))
      return make(l->left, l->data, make(l->right, e, r));
    else
      return make(make(l->left, l->data, l->right->left),
		  l->right->data,
		  make(l->right->right, e, r));
  } else if (hr > hl+1) {
    //@assert hr == hl+2;
    if (height(r->right) >= height(r->left))
      return make(make(l, e, r->left), r->data, r->right);
    else
      return make(make(l, e, r->left->left),
		  r->left->data,
		  make(r->left->right, r->data, r->right));
  } else {
    return make(l, e, r);
  }
}

bool is_bst(bst B) {
  if (B == NULL) return false;
  return is_ordtree(B->root);
}

bst bst_new()
//@ensures is_bst(\result);
{
  bst B = alloc(struct bst);
  B->root = NULL;
  return B;
}

elem tree_search(tree T, key k)
//@requires is_ordtree(T);
//@ensures \result == NULL || compare(elem_key(\result), k) == 0;
{
  if (T == NULL) return NULL;
  int r = compare(k, elem_key(T->data));
  if (r == 0)
    return T->data;
  else if (r < 0)
    return tree_search(T->left, k);
  else //@assert r > 0;
    return tree_search(T->right, k);
}

elem bst_search(bst B, key k)
//@requires is_bst(B);
//@ensures \result == NULL || compare(elem_key(\result), k) == 0;
{
  return tree_search(B->root, k);
}

/* tree_insert(T, e) returns the modified tree
 * this avoids some complications in case T = NULL
 */
tree tree_insert(tree T, elem e)
//@requires is_ordtree(T);
//@ensures is_ordtree(\result);
{
  assert(e != NULL);		/* cannot insert NULL element */
  if (T == NULL) {
    return make(NULL, e, NULL);
  } else {
    int r = compare(elem_key(e), elem_key(T->data));
    if (r < 0)
      return balance(T->left, T->data, tree_insert(T->right, e));
    else if (r == 0)
      return make(T->left, e, T->right);
    else //@assert r > 0;
      return balance(T->left, T->data, tree_insert(T->right, e));
    return T;
  }
}

void bst_insert(bst B, elem e)
//@requires is_bst(B);
//@ensures is_bst(B);
{
  assert(e != NULL);		/* cannot insert NULL element */
  B->root = tree_insert(B->root, e);
  return;
}
