module Rat

import UInt
import Base

/*
  Positive rational numbers with a canonical representation.

  Every positive rational has exactly one `Pos` value, so structural
  equality on `Pos` is equality of rationals and there are no
  non-normalized "junk" values.  The representation is a variant of
  the continued-fraction expansion:

    pint(n)     denotes  1 + n
    pfrac(q, y) denotes  q + 1 / (1 + y)

  `posFrom(a, b)` computes the `Pos` for `a / b` with Euclid's
  algorithm, and `pnum`/`pden` read a `Pos` back as a fraction in
  lowest terms.  The key facts proved here are

    posFrom(pnum(x), pden(x)) = x
    posFrom(a, b) = posFrom(c, d)  ⇔  a * d = c * b   (a, b, c, d > 0)
*/

private union Pos {
  pint(UInt)
  pfrac(UInt, Pos)
}

// `lin(x, a, b)` is `a * pnum(x) + b * pden(x)`, computed with a single
// pass over `x`.
private recursive lin(Pos, UInt, UInt) -> UInt {
  lin(pint(n), a, b) = a * (1 + n) + b
  lin(pfrac(q, y), a, b) = lin(y, a * q + b, a * q + a + b)
}

private fun pnum(x : Pos) { lin(x, 1, 0) }

private fun pden(x : Pos) { lin(x, 0, 1) }

// The `Pos` for `a / b`, by Euclid's algorithm.  Junk inputs
// (`a = 0` or `b = 0`) produce `pint(0)`.
private recfun posFrom(a : UInt, b : UInt) -> Pos
  measure b of UInt
{
  if b = 0 then pint(0)
  else if a % b = 0 then pint(a / b ∸ 1)
  else pfrac(a / b, posFrom(b ∸ a % b, a % b))
}
terminates {
  arbitrary a:UInt, b:UInt
  assume cond: not (b = 0) and not (a % b = 0)
  have b_pos: 0 < b by apply uint_not_zero_pos to conjunct 0 of cond
  conclude a % b < b by apply uint_mod_less_divisor[a, b] to b_pos
}

assert posFrom(6, 4) = pfrac(1, pint(0))
assert pnum(posFrom(6, 4)) = 3
assert pden(posFrom(6, 4)) = 2
assert posFrom(2, 7) = pfrac(0, pfrac(2, pint(0)))
assert pnum(posFrom(2, 7)) = 2
assert pden(posFrom(2, 7)) = 7

lemma pos_lin_linear: all x:Pos. all a:UInt, b:UInt.
  lin(x, a, b) = a * pnum(x) + b * pden(x)
proof
  induction Pos
  case pint(n) {
    arbitrary a:UInt, b:UInt
    expand pnum | pden | lin.
  }
  case pfrac(q, y) assume IH {
    arbitrary a:UInt, b:UInt
    expand pnum | pden | lin
    replace IH | uint_dist_mult_add_right | uint_dist_mult_add | uint_dist_mult_add_right | uint_dist_mult_add
    replace uint_add_commute[b * pnum(y), a * q * pden(y) + a * pden(y)].
  }
end

lemma pos_pnum_pint: all n:UInt. pnum(pint(n)) = 1 + n
proof
  arbitrary n:UInt
  expand pnum | lin.
end

lemma pos_pden_pint: all n:UInt. pden(pint(n)) = 1
proof
  arbitrary n:UInt
  expand pden | lin.
end

lemma pos_pden_pfrac: all q:UInt, y:Pos. pden(pfrac(q, y)) = pnum(y) + pden(y)
proof
  arbitrary q:UInt, y:Pos
  expand pden | lin
  replace pos_lin_linear[y].
end

lemma pos_pnum_pfrac: all q:UInt, y:Pos.
  pnum(pfrac(q, y)) = q * (pnum(y) + pden(y)) + pden(y)
proof
  arbitrary q:UInt, y:Pos
  expand pnum | lin
  replace pos_lin_linear[y]
  replace uint_dist_mult_add_right | uint_dist_mult_add.
end

lemma uint_pos_add_right: all a:UInt, b:UInt. if 0 < b then 0 < a + b
proof
  arbitrary a:UInt, b:UInt
  assume b_pos
  obtain b' where eq: b = 1 + b' from apply uint_positive_add_one to b_pos
  replace eq | uint_add_commute[a, 1 + b'].
end

lemma uint_pos_add_left: all a:UInt, b:UInt. if 0 < a then 0 < a + b
proof
  arbitrary a:UInt, b:UInt
  assume a_pos
  replace uint_add_commute[a, b]
  apply uint_pos_add_right[b, a] to a_pos
end

lemma pos_pnum_pden_pos: all x:Pos. 0 < pnum(x) and 0 < pden(x)
proof
  induction Pos
  case pint(n) {
    replace pos_pnum_pint | pos_pden_pint.
  }
  case pfrac(q, y) assume IH {
    replace pos_pnum_pfrac | pos_pden_pfrac
    have dp: 0 < pnum(y) + pden(y)
      by apply uint_pos_add_left[pnum(y), pden(y)] to conjunct 0 of IH
    (apply uint_pos_add_right[q * (pnum(y) + pden(y)), pden(y)] to conjunct 1 of IH), dp
  }
end

lemma pos_pnum_pos: all x:Pos. 0 < pnum(x)
proof
  arbitrary x:Pos
  conjunct 0 of pos_pnum_pden_pos[x]
end

lemma pos_pden_pos: all x:Pos. 0 < pden(x)
proof
  arbitrary x:Pos
  conjunct 1 of pos_pnum_pden_pos[x]
end

lemma uint_div_small: all n:UInt, m:UInt. if n < m then n / m = 0
proof
  arbitrary n:UInt, m:UInt
  assume lt
  expand operator/
  replace (apply eq_true to lt).
end

lemma uint_less_add_pos: all a:UInt, b:UInt. if 0 < a then b < a + b
proof
  arbitrary a:UInt, b:UInt
  assume a_pos
  replace uint_add_commute[a, b]
  have h: b + 0 < b + a by apply uint_add_both_sides_of_less[b, 0, a] to a_pos
  h
end

lemma pos_from_pnum_pden: all x:Pos. posFrom(pnum(x), pden(x)) = x
proof
  induction Pos
  case pint(n) {
    replace pos_pnum_pint | pos_pden_pint
    expand posFrom
    replace uint_mod_one | uint_div_one.
  }
  case pfrac(q, y) assume IH {
    replace pos_pnum_pfrac | pos_pden_pfrac
    have S_pos: 0 < pnum(y) + pden(y)
      by apply uint_pos_add_left[pnum(y), pden(y)] to pos_pnum_pos[y]
    have D_lt: pden(y) < pnum(y) + pden(y)
      by apply uint_less_add_pos[pnum(y), pden(y)] to pos_pnum_pos[y]
    have mod_eq: (q * (pnum(y) + pden(y)) + pden(y)) % (pnum(y) + pden(y)) = pden(y)
      by transitive (apply uint_mult_add_mod[q, pden(y), pnum(y) + pden(y)] to S_pos)
                    (apply uint_mod_small[pden(y), pnum(y) + pden(y)] to D_lt)
    have div_eq: (q * (pnum(y) + pden(y)) + pden(y)) / (pnum(y) + pden(y)) = q
      by replace (apply uint_div_small[pden(y), pnum(y) + pden(y)] to D_lt) in
         apply uint_mult_add_div[q, pden(y), pnum(y) + pden(y)] to S_pos
    expand posFrom
    replace mod_eq | div_eq
      | apply eq_false to (apply uint_pos_not_zero to S_pos)
      | apply eq_false to (apply uint_pos_not_zero to pos_pden_pos[y])
      | uint_add_commute[pnum(y), pden(y)]
      | IH.
  }
end

lemma uint_div_scale: all k:UInt, a:UInt, b:UInt.
  if 0 < k and 0 < b then (k * a) / (k * b) = a / b
proof
  arbitrary k:UInt, a:UInt, b:UInt
  assume prem
  have k_pos: 0 < k by prem
  have b_pos: 0 < b by prem
  have kb_pos: 0 < k * b
    by apply uint_pos_mult_both_sides_of_less[k, 0, b] to k_pos, b_pos
  have r_lt: k * (a % b) < k * b
    by apply uint_pos_mult_both_sides_of_less[k, a % b, b]
       to k_pos, apply uint_mod_less_divisor[a, b] to b_pos
  have dm: (a / b) * b + a % b = a by apply uint_div_mod[a, b] to b_pos
  have ka_eq: k * a = (a / b) * (k * b) + k * (a % b) by symmetric equations
        (a / b) * (k * b) + k * (a % b)
      = k * (a / b) * b + k * (a % b)     by replace uint_mult_commute[a / b, k].
  ... = # k * ((a / b) * b + a % b) #     by replace uint_dist_mult_add.
  ... = k * a                             by replace dm.
  replace ka_eq
  replace (apply uint_mult_add_div[a / b, k * (a % b), k * b] to kb_pos)
        | (apply uint_div_small[k * (a % b), k * b] to r_lt).
end

lemma uint_mod_scale: all k:UInt, a:UInt, b:UInt.
  if 0 < k and 0 < b then (k * a) % (k * b) = k * (a % b)
proof
  arbitrary k:UInt, a:UInt, b:UInt
  assume prem
  expand operator%
  replace (apply uint_div_scale[k, a, b] to prem)
        | uint_dist_mult_monus[k]
        | uint_mult_commute[a / b, k * b]
        | uint_mult_commute[a / b, b].
end

lemma uint_pos_mult_eq_zero: all k:UInt, x:UInt. if 0 < k then (k * x = 0) = (x = 0)
proof
  arbitrary k:UInt, x:UInt
  assume k_pos
  have k_nz: not (k = 0) by apply uint_pos_not_zero to k_pos
  switch x = 0 {
    case true assume xz { replace xz. }
    case false assume xnz {
      have nkx: not (k * x = 0) by {
        assume kx
        cases apply uint_mult_to_zero[k, x] to kx
        case kz { conclude false by apply k_nz to kz }
        case xz { conclude false by apply xnz to xz }
      }
      apply eq_false to nkx
    }
  }
end

lemma pos_from_scale: all b:UInt, a:UInt, k:UInt.
  if 0 < k then posFrom(k * a, k * b) = posFrom(a, b)
proof
  define P = fun b:UInt { all a:UInt, k:UInt.
                          if 0 < k then posFrom(k * a, k * b) = posFrom(a, b) }
  have X: all j:UInt. (if (all i:UInt. (if i < j then P(i))) then P(j)) by {
    arbitrary j:UInt
    assume IH: all i:UInt. (if i < j then P(i))
    expand P
    arbitrary a:UInt, k:UInt
    assume k_pos
    cases uint_zero_or_positive[j]
    case j_z {
      replace j_z
      expand posFrom.
    }
    case j_pos {
      have kj_pos: 0 < k * j
        by apply uint_pos_mult_both_sides_of_less[k, 0, j] to k_pos, j_pos
      expand posFrom
      replace apply eq_false to (apply uint_pos_not_zero to kj_pos)
            | apply eq_false to (apply uint_pos_not_zero to j_pos)
            | (apply uint_mod_scale[k, a, j] to k_pos, j_pos)
            | (apply uint_div_scale[k, a, j] to k_pos, j_pos)
            | (apply uint_pos_mult_eq_zero[k, a % j] to k_pos)
            | symmetric uint_dist_mult_monus[k, j, a % j]
      have r_lt: a % j < j by apply uint_mod_less_divisor[a, j] to j_pos
      have IH': all a':UInt, k':UInt.
          if 0 < k' then posFrom(k' * a', k' * (a % j)) = posFrom(a', a % j)
        by expand P in apply IH[a % j] to r_lt
      replace apply IH'[j ∸ a % j, k] to k_pos.
    }
  }
  arbitrary b:UInt
  expand P in apply uint_strong_induction[P, b] to X
end

lemma pos_cross_alg: all a:UInt, b:UInt, q:UInt, r:UInt, c:UInt, n:UInt, d:UInt.
  if q * b + r = a and b = r + c and c * d = n * r
  then a * (n + d) = (q * (n + d) + d) * b
proof
  arbitrary a:UInt, b:UInt, q:UInt, r:UInt, c:UInt, n:UInt, d:UInt
  assume prem
  have e1: q * b + r = a by prem
  have e2: b = r + c by prem
  have e3: c * d = n * r by prem
  equations
        a * (n + d)
      = (q * b + r) * (n + d)          by replace symmetric e1.
  ... = q * b * (n + d) + r * (n + d)  by replace uint_dist_mult_add_right[q * b, r, n + d].
  ... = q * (n + d) * b + r * (n + d)  by replace uint_mult_commute[b, n + d].
  ... = q * (n + d) * b + (n * r + d * r)
        by replace uint_dist_mult_add[r, n, d] | uint_mult_commute[r, n] | uint_mult_commute[r, d].
  ... = q * (n + d) * b + (d * c + d * r)
        by replace symmetric e3 | uint_mult_commute[c, d].
  ... = q * (n + d) * b + d * b
        by replace symmetric uint_dist_mult_add[d, c, r] | uint_add_commute[c, r] | symmetric e2.
  ... = (q * (n + d) + d) * b
        by replace symmetric uint_dist_mult_add_right[q * (n + d), d, b].
end

lemma uint_monus_pos: all r:UInt, j:UInt. if r < j then 0 < j ∸ r
proof
  arbitrary r:UInt, j:UInt
  assume r_lt
  have le: 1 + r ≤ j by replace uint_less_is_less_equal in r_lt
  obtain x where jx: j = (1 + r) + x from apply uint_le_exists_monus[1 + r, j] to le
  replace jx | uint_add_commute[1, r].
end

lemma pos_from_cross: all b:UInt, a:UInt.
  if 0 < a and 0 < b then a * pden(posFrom(a, b)) = pnum(posFrom(a, b)) * b
proof
  define P = fun b:UInt { all a:UInt.
    if 0 < a and 0 < b then a * pden(posFrom(a, b)) = pnum(posFrom(a, b)) * b }
  have X: all j:UInt. (if (all i:UInt. (if i < j then P(i))) then P(j)) by {
    arbitrary j:UInt
    assume IH: all i:UInt. (if i < j then P(i))
    expand P
    arbitrary a:UInt
    assume prem
    have a_pos: 0 < a by prem
    have j_pos: 0 < j by prem
    have dm: (a / j) * j + a % j = a by apply uint_div_mod[a, j] to j_pos
    have r_lt: a % j < j by apply uint_mod_less_divisor[a, j] to j_pos
    expand posFrom
    replace apply eq_false to (apply uint_pos_not_zero to j_pos)
    switch a % j = 0 {
      case true assume rz {
        replace pos_pnum_pint | pos_pden_pint
        have Qj: (a / j) * j = a by replace rz in dm
        have Q_pos: 0 < a / j by {
          cases uint_zero_or_positive[a / j]
          case Qz {
            have az: 0 = a by replace Qz in Qj
            conclude false by replace symmetric az in a_pos
          }
          case Qp { Qp }
        }
        replace apply uint_monus_add_identity[a / j, 1] to apply uint_pos_implies_one_le to Q_pos
        symmetric Qj
      }
      case false assume rnz {
        replace pos_pnum_pfrac | pos_pden_pfrac
        have r_pos: 0 < a % j by apply uint_not_zero_pos to rnz
        have c_pos: 0 < j ∸ a % j by apply uint_monus_pos[a % j, j] to r_lt
        have ih: (j ∸ a % j) * pden(posFrom(j ∸ a % j, a % j))
               = pnum(posFrom(j ∸ a % j, a % j)) * (a % j)
          by apply (expand P in apply IH[a % j] to r_lt)[j ∸ a % j] to c_pos, r_pos
        have jeq: j = a % j + (j ∸ a % j)
          by symmetric (apply uint_monus_add_identity[j, a % j]
                        to apply uint_less_implies_less_equal to r_lt)
        apply pos_cross_alg[a, j, a / j, a % j, j ∸ a % j,
                            pnum(posFrom(j ∸ a % j, a % j)), pden(posFrom(j ∸ a % j, a % j))]
        to dm, jeq, ih
      }
    }
  }
  arbitrary b:UInt
  expand P in apply uint_strong_induction[P, b] to X
end

lemma pos_from_cross_eq: all a:UInt, b:UInt, c:UInt, d:UInt.
  if 0 < b and 0 < d and a * d = c * b then posFrom(a, b) = posFrom(c, d)
proof
  arbitrary a:UInt, b:UInt, c:UInt, d:UInt
  assume prem
  have b_pos: 0 < b by prem
  have d_pos: 0 < d by prem
  have eq: a * d = c * b by prem
  equations
        posFrom(a, b)
      = posFrom(d * a, d * b)  by symmetric apply pos_from_scale[b, a, d] to d_pos
  ... = posFrom(b * c, b * d)  by replace uint_mult_commute[d, a] | eq | uint_mult_commute[c, b]
                                        | uint_mult_commute[d, b].
  ... = posFrom(c, d)          by apply pos_from_scale[d, c, b] to b_pos
end

lemma pos_from_eq_cross: all a:UInt, b:UInt, c:UInt, d:UInt.
  if 0 < a and 0 < b and 0 < c and 0 < d and posFrom(a, b) = posFrom(c, d)
  then a * d = c * b
proof
  arbitrary a:UInt, b:UInt, c:UInt, d:UInt
  assume prem
  have e1: a * pden(posFrom(a, b)) = pnum(posFrom(a, b)) * b
    by apply pos_from_cross[b, a] to prem
  have e2: c * pden(posFrom(a, b)) = pnum(posFrom(a, b)) * d
    by replace symmetric (conjunct 4 of prem) in
       apply pos_from_cross[d, c] to prem
  have D_pos: 0 < pden(posFrom(a, b)) by pos_pden_pos[posFrom(a, b)]
  have eq: (a * d) * pden(posFrom(a, b)) = (c * b) * pden(posFrom(a, b)) by
    equations
          (a * d) * pden(posFrom(a, b))
        = d * (a * pden(posFrom(a, b)))   by replace uint_mult_commute[a, d].
    ... = d * pnum(posFrom(a, b)) * b     by replace e1.
    ... = b * (pnum(posFrom(a, b)) * d)   by replace uint_mult_commute[d, pnum(posFrom(a, b)) * b]
                                                   | uint_mult_commute[pnum(posFrom(a, b)), b].
    ... = b * (c * pden(posFrom(a, b)))   by replace symmetric e2.
    ... = (c * b) * pden(posFrom(a, b))   by replace uint_mult_commute[b, c].
  apply uint_pos_mult_right_cancel[pden(posFrom(a, b)), a * d, c * b] to D_pos, eq
end

lemma pos_from_eq_iff: all a:UInt, b:UInt, c:UInt, d:UInt.
  if 0 < a and 0 < b and 0 < c and 0 < d
  then (posFrom(a, b) = posFrom(c, d)) ⇔ (a * d = c * b)
proof
  arbitrary a:UInt, b:UInt, c:UInt, d:UInt
  assume prem
  have fwd: if posFrom(a, b) = posFrom(c, d) then a * d = c * b by {
    assume eq
    apply pos_from_eq_cross[a, b, c, d] to prem, eq
  }
  have bwd: if a * d = c * b then posFrom(a, b) = posFrom(c, d) by {
    assume eq
    apply pos_from_cross_eq[a, b, c, d] to prem, eq
  }
  fwd, bwd
end