module Rat

import Nat
import UInt
import Int
import Base
import RatPos
import RatDefs

/*
  The fraction interface to Rat.

  Every rational is `frac(num(x), den(x))`, and two fractions with
  positive denominators are equal exactly when their cross products
  are equal.  The algebraic laws in the other Rat files are proved by
  reducing them to Int identities through these facts.
*/

theorem rat_den_pos: all x:Rat. 0 < den(x)
proof
  arbitrary x:Rat
  switch x {
    case rzero { expand den. }
    case rpos(p) { expand den  pos_pden_pos[p] }
    case rneg(p) { expand den  pos_pden_pos[p] }
  }
end

theorem rat_frac_num_den: all x:Rat. frac(num(x), den(x)) = x
proof
  arbitrary x:Rat
  switch x {
    case rzero { expand num | den | frac. }
    case rpos(p) {
      expand num | den | frac
      replace apply eq_false to (apply uint_pos_not_zero to pos_pden_pos[p])
            | apply eq_false to (apply uint_pos_not_zero to pos_pnum_pos[p])
            | pos_from_pnum_pden.
    }
    case rneg(p) {
      obtain m where pm: pnum(p) = 1 + m from apply uint_positive_add_one to pos_pnum_pos[p]
      expand num | den | frac
      replace apply eq_false to (apply uint_pos_not_zero to pos_pden_pos[p])
            | pm | neg_pos | symmetric pm | pos_from_pnum_pden.
    }
  }
end

lemma int_neg_negsuc: all a:UInt. - negsuc(a) = pos(1 + a)
proof
  arbitrary a:UInt
  replace symmetric neg_pos[a] | neg_involutive.
end

theorem rat_frac_neg: all n:Int, d:UInt. frac(- n, d) = - frac(n, d)
proof
  arbitrary n:Int, d:UInt
  cases uint_zero_or_positive[d]
  case dz { replace dz  expand frac | operator-. }
  case d_pos {
    have dnz: (d = 0) = false by apply eq_false to (apply uint_pos_not_zero to d_pos)
    switch n {
      case pos(a) {
        cases uint_zero_or_add_one[a]
        case az { replace az  expand frac | operator-  replace dnz. }
        case a_suc {
          obtain a' where aeq: a = 1 + a' from a_suc
          replace aeq | neg_pos
          expand frac | operator-
          replace dnz.
        }
      }
      case negsuc(a) {
        replace int_neg_negsuc
        expand frac | operator-
        replace dnz.
      }
    }
  }
end

lemma rat_frac_scale_pos: all k:UInt, a:UInt, d:UInt.
  if 0 < k then frac(pos(k) * pos(a), k * d) = frac(pos(a), d)
proof
  arbitrary k:UInt, a:UInt, d:UInt
  assume k_pos
  replace mult_pos_pos
  cases uint_zero_or_positive[d]
  case dz { replace dz  expand frac. }
  case d_pos {
    have kd_pos: 0 < k * d
      by apply uint_pos_mult_both_sides_of_less[k, 0, d] to k_pos, d_pos
    expand frac
    replace apply eq_false to (apply uint_pos_not_zero to kd_pos)
          | apply eq_false to (apply uint_pos_not_zero to d_pos)
          | apply uint_pos_mult_eq_zero[k, a] to k_pos
          | apply pos_from_scale[d, a, k] to k_pos.
  }
end

theorem rat_frac_scale: all k:UInt, n:Int, d:UInt.
  if 0 < k then frac(pos(k) * n, k * d) = frac(n, d)
proof
  arbitrary k:UInt, n:Int, d:UInt
  assume k_pos
  switch n {
    case pos(a) { apply rat_frac_scale_pos[k, a, d] to k_pos }
    case negsuc(a) {
      replace symmetric neg_pos[a] | int_neg_mult_right | rat_frac_neg
            | apply rat_frac_scale_pos[k, 1 + a, d] to k_pos.
    }
  }
end

theorem rat_frac_cross_eq: all n:Int, d:UInt, m:Int, e:UInt.
  if 0 < d and 0 < e and n * pos(e) = m * pos(d) then frac(n, d) = frac(m, e)
proof
  arbitrary n:Int, d:UInt, m:Int, e:UInt
  assume prem
  have d_pos: 0 < d by prem
  have e_pos: 0 < e by prem
  have eq: n * pos(e) = m * pos(d) by prem
  equations
        frac(n, d)
      = frac(pos(e) * n, e * d)  by symmetric apply rat_frac_scale[e, n, d] to e_pos
  ... = frac(pos(d) * m, d * e)  by replace int_mult_commute[pos(e), n] | eq
                                          | int_mult_commute[m, pos(d)] | uint_mult_commute[e, d].
  ... = frac(m, e)               by apply rat_frac_scale[d, m, e] to d_pos
end

theorem rat_frac_cross: all n:Int, d:UInt.
  if 0 < d then num(frac(n, d)) * pos(d) = n * pos(den(frac(n, d)))
proof
  arbitrary n:Int, d:UInt
  assume d_pos
  have dnz: (d = 0) = false by apply eq_false to (apply uint_pos_not_zero to d_pos)
  switch n {
    case pos(a) {
      cases uint_zero_or_positive[a]
      case az { replace az  expand frac  replace dnz  expand num. }
      case a_pos {
        expand frac
        replace dnz | apply eq_false to (apply uint_pos_not_zero to a_pos)
        expand num | den
        replace mult_pos_pos | apply pos_from_cross[d, a] to a_pos, d_pos.
      }
    }
    case negsuc(a) {
      expand frac
      replace dnz
      expand num | den
      replace symmetric neg_pos[a]
            | symmetric dist_neg_mult[pos(pnum(posFrom(1 + a, d))), pos(d)]
            | symmetric dist_neg_mult[pos(1 + a), pos(pden(posFrom(1 + a, d)))]
            | mult_pos_pos
            | apply pos_from_cross[d, 1 + a] to uint_zero_less_one_add[a], d_pos.
    }
  }
end

theorem rat_frac_eq_cross: all n:Int, d:UInt, m:Int, e:UInt.
  if 0 < d and 0 < e and frac(n, d) = frac(m, e) then n * pos(e) = m * pos(d)
proof
  arbitrary n:Int, d:UInt, m:Int, e:UInt
  assume prem
  have d_pos: 0 < d by prem
  have e_pos: 0 < e by prem
  have eq: frac(n, d) = frac(m, e) by prem
  have c1: num(frac(n, d)) * pos(d) = n * pos(den(frac(n, d)))
    by apply rat_frac_cross[n, d] to d_pos
  have c2: num(frac(n, d)) * pos(e) = m * pos(den(frac(n, d)))
    by replace symmetric eq in apply rat_frac_cross[m, e] to e_pos
  have Dnz: not (pos(den(frac(n, d))) = +0) by {
    assume z
    have dz: den(frac(n, d)) = 0 by apply int_pos_injective to z
    conclude false by apply (apply uint_pos_not_zero to rat_den_pos[frac(n, d)]) to dz
  }
  have prod: (n * pos(e)) * pos(den(frac(n, d))) = (m * pos(d)) * pos(den(frac(n, d))) by
    equations
          (n * pos(e)) * pos(den(frac(n, d)))
        = pos(e) * (n * pos(den(frac(n, d))))  by replace int_mult_commute[n, pos(e)].
    ... = pos(e) * num(frac(n, d)) * pos(d)    by replace symmetric c1.
    ... = pos(d) * (num(frac(n, d)) * pos(e))  by replace int_mult_commute[pos(e), num(frac(n, d)) * pos(d)]
                                                     | int_mult_commute[num(frac(n, d)), pos(d)].
    ... = pos(d) * (m * pos(den(frac(n, d))))  by replace c2.
    ... = (m * pos(d)) * pos(den(frac(n, d)))  by replace int_mult_commute[pos(d), m].
  apply int_mult_right_cancel[pos(den(frac(n, d))), n * pos(e), m * pos(d)] to Dnz, prod
end

lemma rat_add_frac_alg: all a:Int, b:Int, c:Int, g:Int, n:Int, m:Int, d:Int, e:Int.
  if a * d = n * g and c * e = m * b
  then (a * b + c * g) * (d * e) = (n * e + m * d) * (g * b)
proof
  arbitrary a:Int, b:Int, c:Int, g:Int, n:Int, m:Int, d:Int, e:Int
  assume prem
  have e1: a * d = n * g by prem
  have e2: c * e = m * b by prem
  equations
        (a * b + c * g) * (d * e)
      = a * b * d * e + c * g * d * e    by replace int_dist_mult_add_right.
  ... = (a * d) * b * e + (c * e) * g * d
        by replace int_mult_commute[b, d] | int_mult_commute[g * d, e].
  ... = n * g * b * e + m * b * g * d    by replace e1 | e2.
  ... = n * e * (g * b) + m * d * (g * b)
        by replace int_mult_commute[g * b, e] | int_mult_commute[b * g, d]
                 | int_mult_commute[b, g].
  ... = # (n * e + m * d) * (g * b) #    by replace int_dist_mult_add_right.
end

theorem rat_add_frac: all n:Int, d:UInt, m:Int, e:UInt.
  if 0 < d and 0 < e then frac(n, d) + frac(m, e) = frac(n * pos(e) + m * pos(d), d * e)
proof
  arbitrary n:Int, d:UInt, m:Int, e:UInt
  assume prem
  have d_pos: 0 < d by prem
  have e_pos: 0 < e by prem
  have c1: num(frac(n, d)) * pos(d) = n * pos(den(frac(n, d)))
    by apply rat_frac_cross[n, d] to d_pos
  have c2: num(frac(m, e)) * pos(e) = m * pos(den(frac(m, e)))
    by apply rat_frac_cross[m, e] to e_pos
  have de_pos: 0 < d * e
    by apply uint_pos_mult_both_sides_of_less[d, 0, e] to d_pos, e_pos
  have DE_pos: 0 < den(frac(n, d)) * den(frac(m, e))
    by apply uint_pos_mult_both_sides_of_less[den(frac(n, d)), 0, den(frac(m, e))]
       to rat_den_pos[frac(n, d)], rat_den_pos[frac(m, e)]
  expand operator+
  apply rat_frac_cross_eq[num(frac(n, d)) * pos(den(frac(m, e))) + num(frac(m, e)) * pos(den(frac(n, d))),
                          den(frac(n, d)) * den(frac(m, e)),
                          n * pos(e) + m * pos(d), d * e]
  to DE_pos, de_pos,
     (replace mult_pos_pos[d, e] | mult_pos_pos[den(frac(n, d)), den(frac(m, e))]
      in apply rat_add_frac_alg[num(frac(n, d)), pos(den(frac(m, e))), num(frac(m, e)),
                                pos(den(frac(n, d))), n, m, pos(d), pos(e)] to c1, c2)
end

theorem rat_mult_frac: all n:Int, d:UInt, m:Int, e:UInt.
  if 0 < d and 0 < e then frac(n, d) * frac(m, e) = frac(n * m, d * e)
proof
  arbitrary n:Int, d:UInt, m:Int, e:UInt
  assume prem
  have d_pos: 0 < d by prem
  have e_pos: 0 < e by prem
  have c1: num(frac(n, d)) * pos(d) = n * pos(den(frac(n, d)))
    by apply rat_frac_cross[n, d] to d_pos
  have c2: num(frac(m, e)) * pos(e) = m * pos(den(frac(m, e)))
    by apply rat_frac_cross[m, e] to e_pos
  have de_pos: 0 < d * e
    by apply uint_pos_mult_both_sides_of_less[d, 0, e] to d_pos, e_pos
  have DE_pos: 0 < den(frac(n, d)) * den(frac(m, e))
    by apply uint_pos_mult_both_sides_of_less[den(frac(n, d)), 0, den(frac(m, e))]
       to rat_den_pos[frac(n, d)], rat_den_pos[frac(m, e)]
  expand operator*
  have eq: (num(frac(n, d)) * num(frac(m, e))) * pos(d * e)
         = (n * m) * pos(den(frac(n, d)) * den(frac(m, e))) by
    equations
          (num(frac(n, d)) * num(frac(m, e))) * pos(d * e)
        = (num(frac(n, d)) * pos(d)) * (num(frac(m, e)) * pos(e))
          by replace symmetric mult_pos_pos[d, e] | int_mult_commute[num(frac(m, e)), pos(d)].
    ... = n * pos(den(frac(n, d))) * m * pos(den(frac(m, e)))  by replace c1 | c2.
    ... = (n * m) * pos(den(frac(n, d)) * den(frac(m, e)))
          by replace int_mult_commute[pos(den(frac(n, d))), m] | mult_pos_pos.
  apply rat_frac_cross_eq[num(frac(n, d)) * num(frac(m, e)), den(frac(n, d)) * den(frac(m, e)),
                          n * m, d * e]
  to DE_pos, de_pos, eq
end

lemma int_pos_mult_le_iff: all k:Int, x:Int, y:Int.
  if +0 < k then (x * k ≤ y * k) = (x ≤ y)
proof
  arbitrary k:Int, x:Int, y:Int
  assume k_pos
  have fwd: if x * k ≤ y * k then x ≤ y by {
    assume h
    apply int_pos_mult_right_cancel_le[k, x, y] to k_pos, h
  }
  have bwd: if x ≤ y then x * k ≤ y * k by {
    assume h
    apply int_nonneg_mult_mono_le_right[k, x, y]
      to (apply int_less_implies_less_equal[+0, k] to k_pos), h
  }
  apply iff_equal to fwd, bwd
end

lemma int_pos_mult_less_iff: all k:Int, x:Int, y:Int.
  if +0 < k then (x * k < y * k) = (x < y)
proof
  arbitrary k:Int, x:Int, y:Int
  assume k_pos
  have fwd: if x * k < y * k then x < y by {
    assume h
    apply int_pos_mult_right_cancel_less[k, x, y] to k_pos, h
  }
  have bwd: if x < y then x * k < y * k by {
    assume h
    apply int_pos_mult_mono_less_right[k, x, y] to k_pos, h
  }
  apply iff_equal to fwd, bwd
end

lemma rat_cmp_frac_alg: all a:Int, b:Int, g:Int, n:Int, d:Int, e:Int.
  if a * d = n * g then (a * b) * (d * e) = (n * e) * (g * b)
proof
  arbitrary a:Int, b:Int, g:Int, n:Int, d:Int, e:Int
  assume eq
  equations
        (a * b) * (d * e)
      = (a * d) * b * e  by replace int_mult_commute[b, d].
  ... = n * g * b * e    by replace eq.
  ... = (n * e) * (g * b) by replace int_mult_commute[g * b, e].
end

lemma rat_cmp_frac: all n:Int, d:UInt, m:Int, e:UInt.
  if 0 < d and 0 < e
  then (num(frac(n, d)) * pos(den(frac(m, e)))) * (pos(d) * pos(e))
         = (n * pos(e)) * (pos(den(frac(n, d))) * pos(den(frac(m, e))))
   and (num(frac(m, e)) * pos(den(frac(n, d)))) * (pos(d) * pos(e))
         = (m * pos(d)) * (pos(den(frac(n, d))) * pos(den(frac(m, e))))
proof
  arbitrary n:Int, d:UInt, m:Int, e:UInt
  assume prem
  have c1: num(frac(n, d)) * pos(d) = n * pos(den(frac(n, d)))
    by apply rat_frac_cross[n, d] to prem
  have c2: num(frac(m, e)) * pos(e) = m * pos(den(frac(m, e)))
    by apply rat_frac_cross[m, e] to prem
  have A: (num(frac(n, d)) * pos(den(frac(m, e)))) * (pos(d) * pos(e))
         = (n * pos(e)) * (pos(den(frac(n, d))) * pos(den(frac(m, e))))
    by apply rat_cmp_frac_alg[num(frac(n, d)), pos(den(frac(m, e))), pos(den(frac(n, d))),
                              n, pos(d), pos(e)] to c1
  have B: (num(frac(m, e)) * pos(den(frac(n, d)))) * (pos(e) * pos(d))
         = (m * pos(d)) * (pos(den(frac(m, e))) * pos(den(frac(n, d))))
    by apply rat_cmp_frac_alg[num(frac(m, e)), pos(den(frac(n, d))), pos(den(frac(m, e))),
                              m, pos(e), pos(d)] to c2
  A, (replace int_mult_commute[pos(e), pos(d)]
        | int_mult_commute[pos(den(frac(m, e))), pos(den(frac(n, d)))] in B)
end

lemma rat_den_int_pos: all x:Rat. +0 < pos(den(x))
proof
  arbitrary x:Rat
  rat_den_pos[x]
end

theorem rat_le_frac: all n:Int, d:UInt, m:Int, e:UInt.
  if 0 < d and 0 < e then (frac(n, d) ≤ frac(m, e)) = (n * pos(e) ≤ m * pos(d))
proof
  arbitrary n:Int, d:UInt, m:Int, e:UInt
  assume prem
  have de_pos: +0 < pos(d) * pos(e) by {
    replace mult_pos_pos
    apply uint_pos_mult_both_sides_of_less[d, 0, e] to prem
  }
  have DE_pos: +0 < pos(den(frac(n, d))) * pos(den(frac(m, e))) by {
    replace mult_pos_pos
    apply uint_pos_mult_both_sides_of_less[den(frac(n, d)), 0, den(frac(m, e))]
      to rat_den_pos[frac(n, d)], rat_den_pos[frac(m, e)]
  }
  expand operator≤
  replace symmetric apply int_pos_mult_le_iff[pos(d) * pos(e),
                     num(frac(n, d)) * pos(den(frac(m, e))),
                     num(frac(m, e)) * pos(den(frac(n, d)))] to de_pos
        | conjunct 0 of apply rat_cmp_frac[n, d, m, e] to prem
        | conjunct 1 of apply rat_cmp_frac[n, d, m, e] to prem
        | apply int_pos_mult_le_iff[pos(den(frac(n, d))) * pos(den(frac(m, e))),
                                    n * pos(e), m * pos(d)] to DE_pos.
end

theorem rat_less_frac: all n:Int, d:UInt, m:Int, e:UInt.
  if 0 < d and 0 < e then (frac(n, d) < frac(m, e)) = (n * pos(e) < m * pos(d))
proof
  arbitrary n:Int, d:UInt, m:Int, e:UInt
  assume prem
  have de_pos: +0 < pos(d) * pos(e) by {
    replace mult_pos_pos
    apply uint_pos_mult_both_sides_of_less[d, 0, e] to prem
  }
  have DE_pos: +0 < pos(den(frac(n, d))) * pos(den(frac(m, e))) by {
    replace mult_pos_pos
    apply uint_pos_mult_both_sides_of_less[den(frac(n, d)), 0, den(frac(m, e))]
      to rat_den_pos[frac(n, d)], rat_den_pos[frac(m, e)]
  }
  expand operator<
  replace symmetric apply int_pos_mult_less_iff[pos(d) * pos(e),
                     num(frac(n, d)) * pos(den(frac(m, e))),
                     num(frac(m, e)) * pos(den(frac(n, d)))] to de_pos
        | conjunct 0 of apply rat_cmp_frac[n, d, m, e] to prem
        | conjunct 1 of apply rat_cmp_frac[n, d, m, e] to prem
        | apply int_pos_mult_less_iff[pos(den(frac(n, d))) * pos(den(frac(m, e))),
                                    n * pos(e), m * pos(d)] to DE_pos.
end

theorem rat_frac_rep: all x:Rat. some n:Int, d:UInt. 0 < d and x = frac(n, d)
proof
  arbitrary x:Rat
  choose num(x), den(x)
  rat_den_pos[x], symmetric rat_frac_num_den[x]
end

theorem rat_frac_eq_iff: all n:Int, d:UInt, m:Int, e:UInt.
  if 0 < d and 0 < e then (frac(n, d) = frac(m, e)) = (n * pos(e) = m * pos(d))
proof
  arbitrary n:Int, d:UInt, m:Int, e:UInt
  assume prem
  have fwd: if frac(n, d) = frac(m, e) then n * pos(e) = m * pos(d) by {
    assume eq
    apply rat_frac_eq_cross[n, d, m, e] to prem, eq
  }
  have bwd: if n * pos(e) = m * pos(d) then frac(n, d) = frac(m, e) by {
    assume eq
    apply rat_frac_cross_eq[n, d, m, e] to prem, eq
  }
  apply iff_equal to fwd, bwd
end

lemma uint_mult_pos: all a:UInt, b:UInt. if 0 < a and 0 < b then 0 < a * b
proof
  arbitrary a:UInt, b:UInt
  assume prem
  apply uint_pos_mult_both_sides_of_less[a, 0, b] to prem
end