Rat.sqrt

-- THE RATIONAL APPROXIMANT to √u, and the four facts about it that
-- everything downstream uses.
--
--   qSqrtAt u D  =  isqrt ⌊u⁺·(D+1)²⌋ / (D+1)
--
-- Written on the POSITIVE PART u⁺ = max u 0, so it is total: no
-- hypothesis is needed to form it, and on nonnegative u it is what it
-- should be. That matters because a positivity witness on ℝ lives in
-- a squash and cannot be eliminated into data (B-15) — a square root
-- that demanded one could not be defined at all.
--
-- The interface is deliberately small. Downstream never unfolds
-- qSqrtAt; it uses
--
--   leQFlBelow : ⌊w⌋ ≤ w                  (ratFloor)
--   leQFlAbove : w ≤ ⌊w⌋ + 1              (ratCeil)
--   isqrtLe    : K² ≤ ⌊w⌋                 (natSqrt)
--   isqrtUpper : ⌊w⌋ < (K+1)²             (natSqrt)
--
-- and the two bridges below that turn a ℕ comparison into a ℚ one.
-- ===== the approximant =====

import Natural (+, *, plusZeroId, zeroPlusId, plusComm, multZeroId, multSucId, multAssoc)
import Natural.order (≤, leTrans, leMultMonoR, lePlusMonoR)
import Natural.sqrt (isqrt, isqrtLe, isqrtUpper, sqrtStepN)
import Int (Int, intOne)
import Int.mul (*, intMulOneL)
import Rat.frac (NZ, nzOne, nzPos, nzMul, Rat, mkRat)
import Rat (Q, qcls, +, *, qNeg, qZero, qOne, qAddZeroR, qMulCls, ratMul, qMulAssoc, qMulComm, qMulOneR, qDistribR)
import Rat.order (≤, leQTrans, leQMulMono)
import Rat.bound (Bnd, leQAdd)
import Natural.more (multLeftComm)
import Rat.half (leQZeroInvNat)
import Rat.abs (leQZeroMul)
import Rat.nat (qOfNat, leQOfNat, leNOfQ, qOfNatAdd, qOfNatMul, qMulSucNat)
import Rat.ceil (qFloor, qNatBound, leQBound, leQZeroNat)
import Rat.floor (leQFloor)
import Real (qInvNat, rBound)
import Real.mul (mulIdx, qMulInvIdx)
import Core.equality (trans, sym, cong, transport, transportP)

sqNat : ℕ → ℕ
sqNat = λD. S D * S D

qScaled : Q → ℕ → Q
qScaled = λu D. u * qOfNat (sqNat D)

leQZeroScaled : {u : Q} (D : ℕ) → qZero ≤ u → qZero ≤ qScaled u D using (qScaled.eq)
leQZeroScaled = λu D nn. leQZeroMul nn (leQZeroNat (sqNat D))

-- ⌊u⁺·(D+1)²⌋, the natural the root is taken of
qFl : Q → ℕ → ℕ
qFl = λu D. qFloor (qScaled u D)

qSqrtNum : Q → ℕ → ℕ
qSqrtNum = λu D. isqrt (qFl u D)

qSqrtAt : Q → ℕ → Q
qSqrtAt = λu D. qOfNat (qSqrtNum u D) * qInvNat D

-- ===== the four facts =====
leQFlBelow : {u : Q} (D : ℕ) → qZero ≤ u → qOfNat (qFl u D) ≤ qScaled u D using (qFl.eq)
leQFlBelow = λu D nn. leQFloor (leQZeroScaled D nn)

leQFlAbove : (u : Q) (D : ℕ) → qScaled u D ≤ qOfNat (S (qFl u D))
  using (qFl.eq, qScaled.eq, Rat.ceil.qNatBound.eq)
leQFlAbove = λu D. leQBound (qScaled u D)

sqrtNumLe : (u : Q) (D : ℕ) → qSqrtNum u D * qSqrtNum u D ≤ qFl u D using (qSqrtNum.eq)
sqrtNumLe = λu D. isqrtLe (qFl u D)

sqrtNumUpper : (u : Q) (D : ℕ) → S (qFl u D) ≤ S (qSqrtNum u D) * S (qSqrtNum u D)
  using (qSqrtNum.eq)
sqrtNumUpper = λu D. isqrtUpper (qFl u D)

-- ===== the two bridges =====
-- (D+1)·(1/(D+1)) is one, on the nose
qMulInvNat : {D : ℕ} → qOfNat (S D) * qInvNat D ≡ qOne
  using (Int.intOne.eq,
    Rat.frac.mkRat.eq,
    Rat.frac.nzOne.eq,
    Rat.frac.nzPos.eq,
    Rat.frac.ratOfInt.eq,
    Rat.frac.ratOne.eq,
    Rat.qOne.eq,
    Rat.qcls.eq,
    Real.qInvNat.eq,
    Real.mul.mulIdx.eq)
qMulInvNat =
  λD. transportP
    λw. qOfNat (S D) * qInvNat w ≡ qOne
    trans
      mulIdx D Z
      _
      _
      cong
        λw. ℕ
        λw. Z + w
        trans
          _
          _
          _
          multSucId D Z
          trans _ _ _ (cong (λw. ℕ) (λw. D + w) (multZeroId D)) (plusZeroId D)
      zeroPlusId D
    transportP (λw. qOfNat (S D) * qInvNat (mulIdx D Z) ≡ w) {qInvNat Z} {qOne} ⋆ qMulInvIdx

-- 1/((c+1)(m+1)) is 1/(c+1) times 1/(m+1). nzMul on two positive
-- denominators lands on c*m + c + m, which is mulIdx c m written the
-- other way round
mulIdxNz : {c m : ℕ} → c * m + c + m ≡ mulIdx c m using (Real.mul.mulIdx.eq)
mulIdxNz =
  λc m. c * m + c + m
    ≡⟨ plusComm m (c * m + c) ⟩ m + (c * m + c)
    ≡⟨ cong (λw. ℕ) (λw. m + w) (plusComm c (c * m)) ⟩ m + (c + c * m)
    ≡⟨ cong (λw. ℕ) (λw. m + w) (sym _ _ (multSucId c m)) ⟩ m + c * S m

nzPosMul : (c m : ℕ) → nzMul (nzPos c) (nzPos m) ≡ nzPos (mulIdx c m)
  using (Rat.frac.nzMul.eq, Rat.frac.nzPos.eq)
nzPosMul = λc m. cong (λw. NZ) (λw. nzPos w) {c * m + c + m} {mulIdx c m} mulIdxNz

ratMulInv : (c m : ℕ)
  → ratMul (mkRat intOne (nzPos c)) (mkRat intOne (nzPos m))
    ≡ mkRat (intOne * intOne) (nzMul (nzPos c) (nzPos m))
  using (Rat.ratMul.eq, Rat.frac.mkRat.eq, Rat.frac.num.eq, Rat.frac.den.eq)
ratMulInv = λc m. ⋆

qInvNatCls : (k : ℕ) → qInvNat k ≡ qcls (mkRat intOne (nzPos k)) using (Real.qInvNat.eq)
qInvNatCls = λk. ⋆

qInvNatMulIdx : {c m : ℕ} → qInvNat (mulIdx c m) ≡ qInvNat c * qInvNat m
qInvNatMulIdx =
  λc m. qInvNat (mulIdx c m)
    ≡⟨ qInvNatCls (mulIdx c m) ⟩ qcls (mkRat intOne (nzPos (mulIdx c m)))
    ≡⟨ cong (λw. Q) (λw. qcls (mkRat intOne w)) (sym _ _ (nzPosMul c m)) ⟩
      qcls (mkRat intOne (nzMul (nzPos c) (nzPos m)))
    ≡⟨ cong (λw. Q) (λw. qcls (mkRat w (nzMul (nzPos c) (nzPos m)))) (sym _ _ (intMulOneL intOne)) ⟩
      qcls (mkRat (intOne * intOne) (nzMul (nzPos c) (nzPos m)))
    ≡⟨ cong (λw. Q) (λw. qcls w) (sym _ _ (ratMulInv c m)) ⟩
      qcls (ratMul (mkRat intOne (nzPos c)) (mkRat intOne (nzPos m)))
    ≡⟨ sym _ _ (qMulCls (mkRat intOne (nzPos c)) (mkRat intOne (nzPos m))) ⟩
      qcls (mkRat intOne (nzPos c)) * qcls (mkRat intOne (nzPos m))
    ≡⟨ cong (λw. Q) (λw. w * qcls (mkRat intOne (nzPos m))) (sym _ _ (qInvNatCls c)) ⟩
      qInvNat c * qcls (mkRat intOne (nzPos m))
    ≡⟨ cong (λw. Q) (λw. qInvNat c * w) (sym _ _ (qInvNatCls m)) ⟩ qInvNat c * qInvNat m

-- ===== comparing two approximants =====
--
-- The heart. Two approximants of DIFFERENT arguments at DIFFERENT
-- depths are compared by cross-multiplying to naturals, where
-- natSqrt's leOfSqLe applies. The rational work is only to produce
-- the ℕ inequality
--
--   ⌊u⁺(Dm+1)²⌋·(Dn+1)² ≤ (⌊v⁺(Dn+1)²⌋+1)·(Dm+1)² + T
--
-- from u⁺ ≤ v⁺ + d and d·(Dm+1)²(Dn+1)² ≤ T. Everything in it is a
-- chain of ≤; no strict order and no division appear.
qMulSwap3 : (x y z : Q) → x * y * z ≡ x * z * y
qMulSwap3 =
  λx y z. x * y * z
    ≡⟨ qMulAssoc x y z ⟩ x * (y * z)
    ≡⟨ cong (λw. Q) (λw. x * w) (qMulComm y z) ⟩ x * (z * y)
    ≡⟨ sym _ _ (qMulAssoc x z y) ⟩ x * z * y

qDistrib3 : (x y z w : Q) → (x + y) * z * w ≡ x * z * w + y * z * w
qDistrib3 =
  λx y z w. (x + y) * z * w
    ≡⟨ cong (λt. Q) (λt. t * w) (qDistribR z x y) ⟩ (x * z + y * z) * w
    ≡⟨ qDistribR w (x * z) (y * z) ⟩ x * z * w + y * z * w

qNatSumEq : (a b T : ℕ) → qOfNat a * qOfNat b + qOfNat T ≡ qOfNat (a * b + T)
qNatSumEq =
  λa b T. qOfNat a * qOfNat b + qOfNat T
    ≡⟨ cong (λw. Q) (λw. w + qOfNat T) (qOfNatMul a b) ⟩ qOfNat (a * b) + qOfNat T
    ≡⟨ qOfNatAdd (a * b) T ⟩ qOfNat (a * b + T)

flCompare : {u v : Q}
  {Dm Dn T : ℕ}
  (d : Q)
  → qZero ≤ u
    → u ≤ v + d
      → d * qOfNat (sqNat Dm) * qOfNat (sqNat Dn) ≤ qOfNat T
        → qFl u Dm * sqNat Dn ≤ S (qFl v Dn) * sqNat Dm + T
  using (qScaled.eq)
flCompare =
  λu v Dm Dn T d nnu h1 h2. leNOfQ
    transport
      λw. w ≤ qOfNat (S (qFl v Dn) * sqNat Dm + T)
      qOfNatMul (qFl u Dm) (sqNat Dn)
      transport
        λw. qOfNat (qFl u Dm) * qOfNat (sqNat Dn) ≤ w
        qNatSumEq (S (qFl v Dn)) (sqNat Dm) T
        leQTrans
          _
          _
          _
          leQMulMono _ (u * qOfNat (sqNat Dm)) (leQFlBelow Dm nnu) (leQZeroNat (sqNat Dn))
          leQTrans
            _
            _
            _
            leQMulMono _ _ (leQMulMono _ _ h1 (leQZeroNat (sqNat Dm))) (leQZeroNat (sqNat Dn))
            transport
              λw. w ≤ qOfNat (S (qFl v Dn)) * qOfNat (sqNat Dm) + qOfNat T
              sym _ _ (qDistrib3 v d (qOfNat (sqNat Dm)) (qOfNat (sqNat Dn)))
              leQAdd
                _
                _
                _
                _
                transport
                  λw. w ≤ qOfNat (S (qFl v Dn)) * qOfNat (sqNat Dm)
                  sym _ _ (qMulSwap3 v (qOfNat (sqNat Dm)) (qOfNat (sqNat Dn)))
                  leQMulMono (v * qOfNat (sqNat Dn)) _ (leQFlAbove v Dn) (leQZeroNat (sqNat Dm))
                h2

-- ===== ...and back to a comparison of approximants =====
-- (a·b)² = a²·b², so cross-multiplying commutes with squaring
mulSq : (a b : ℕ) → a * b * (a * b) ≡ a * a * (b * b)
mulSq =
  λa b. a * b * (a * b)
    ≡⟨ multAssoc a b (a * b) ⟩ a * (b * (a * b))
    ≡⟨ cong (λw. ℕ) (λw. a * w) (multLeftComm b a b) ⟩ a * (a * (b * b))
    ≡⟨ sym _ _ (multAssoc a a (b * b)) ⟩ a * a * (b * b)

sqrtCompareN : {u v : Q}
  {Dm Dn T : ℕ}
  (d : Q)
  → qZero ≤ u
    → u ≤ v + d
      → d * qOfNat (sqNat Dm) * qOfNat (sqNat Dn) ≤ qOfNat T
        → qSqrtNum u Dm * S Dn ≤ S (qSqrtNum v Dn) * S Dm + S (isqrt T)
  using (sqNat.eq)
sqrtCompareN =
  λu v Dm Dn T d nnu h1 h2. sqrtStepN
    transport
      λw. w ≤ S (qSqrtNum v Dn) * S Dm * (S (qSqrtNum v Dn) * S Dm) + T
      sym _ _ (mulSq (qSqrtNum u Dm) (S Dn))
      transport
        λw. qSqrtNum u Dm * qSqrtNum u Dm * (S Dn * S Dn) ≤ w + T
        sym _ _ (mulSq (S (qSqrtNum v Dn)) (S Dm))
        leTrans
          qSqrtNum u Dm * qSqrtNum u Dm * (S Dn * S Dn)
          qFl u Dm * sqNat Dn
          _
          leMultMonoR (sqNat Dn) (sqrtNumLe u Dm)
          leTrans
            _
            S (qFl v Dn) * sqNat Dm + T
            S (qSqrtNum v Dn) * S (qSqrtNum v Dn) * (S Dm * S Dm) + T
            flCompare _ nnu h1 h2
            lePlusMonoR
              S (qSqrtNum v Dn) * S (qSqrtNum v Dn) * (S Dm * S Dm)
              T
              leMultMonoR (sqNat Dm) (sqrtNumUpper v Dn)

-- ===== ...and back to a comparison of approximants =====
qMulLeftComm : (x y z : Q) → x * (y * z) ≡ y * (x * z)
qMulLeftComm =
  λx y z. x * (y * z)
    ≡⟨ sym _ _ (qMulAssoc x y z) ⟩ x * y * z
    ≡⟨ cong (λw. Q) (λw. w * z) (qMulComm x y) ⟩ y * x * z
    ≡⟨ qMulAssoc y x z ⟩ y * (x * z)

-- the denominator cancels its own numeral, exactly
cancelDen : {A : Q} {D : ℕ} {e : Q} → A * qOfNat (S D) * (e * qInvNat D) ≡ A * e
cancelDen =
  λA D e. A * qOfNat (S D) * (e * qInvNat D)
    ≡⟨ qMulAssoc A (qOfNat (S D)) (e * qInvNat D) ⟩ A * (qOfNat (S D) * (e * qInvNat D))
    ≡⟨ cong (λw. Q) (λw. A * w) (qMulLeftComm (qOfNat (S D)) e (qInvNat D)) ⟩
      A * (e * (qOfNat (S D) * qInvNat D))
    ≡⟨ cong (λw. Q) (λw. A * (e * w)) {qOfNat (S D) * qInvNat D} {qOne} qMulInvNat ⟩ A * (e * qOne)
    ≡⟨ cong (λw. Q) (λw. A * w) (qMulOneR e) ⟩ A * e

-- a ℕ inequality between cross-multiplied numerators IS a ℚ
-- inequality between the fractions
crossToQ : {a b e Dm Dn : ℕ}
  → a * S Dn ≤ b * S Dm + e
    → qOfNat a * qInvNat Dm ≤ qOfNat b * qInvNat Dn + qOfNat e * (qInvNat Dm * qInvNat Dn)
crossToQ =
  λa b e Dm Dn le. transport
    λw. w ≤ qOfNat b * qInvNat Dn + qOfNat e * (qInvNat Dm * qInvNat Dn)
    {qOfNat a * qOfNat (S Dn) * (qInvNat Dm * qInvNat Dn)}
    {qOfNat a * qInvNat Dm}
    cancelDen
    transport
      {Q}
      λw. qOfNat a * qOfNat (S Dn) * (qInvNat Dm * qInvNat Dn) ≤ w
      cong
        λw. Q
        λw. w + qOfNat e * (qInvNat Dm * qInvNat Dn)
        trans
          _
          _
          qOfNat b * qInvNat Dn
          cong (λw. Q) (λw. qOfNat b * qOfNat (S Dm) * w) (qMulComm (qInvNat Dm) (qInvNat Dn))
          cancelDen
      transport
        λw. qOfNat a * qOfNat (S Dn) * (qInvNat Dm * qInvNat Dn) ≤ w
        qDistribR (qInvNat Dm * qInvNat Dn) (qOfNat b * qOfNat (S Dm)) (qOfNat e)
        leQMulMono
          _
          _
          transport
            λw. w ≤ qOfNat b * qOfNat (S Dm) + qOfNat e
            sym _ _ (qOfNatMul a (S Dn))
            transport (λw. qOfNat (a * S Dn) ≤ w) (sym _ _ (qNatSumEq b (S Dm) e)) (leQOfNat le)
          leQZeroMul (leQZeroInvNat Dm) (leQZeroInvNat Dn)

-- ===== THE comparison, in ℚ =====
--
--   √u at depth Dm  ≤  √v at depth Dn  +  1/(Dn+1)  +  (isqrt T + 1)/((Dm+1)(Dn+1))
--
-- given u⁺ ≤ v⁺ + d and d·(Dm+1)²(Dn+1)² ≤ T. The first slack term is
-- the approximation error at n; the second is √ of the gap d, and it
-- is where "√ is not Lipschitz" is paid for.
leQSqrtStep : {u v : Q}
  {Dm Dn T : ℕ}
  (d : Q)
  → qZero ≤ u
    → u ≤ v + d
      → d * qOfNat (sqNat Dm) * qOfNat (sqNat Dn) ≤ qOfNat T
        → qSqrtAt u Dm
          ≤ qSqrtAt v Dn + qInvNat Dn + qOfNat (S (isqrt T)) * (qInvNat Dm * qInvNat Dn)
  using (qSqrtAt.eq)
leQSqrtStep =
  λu v Dm Dn T d nnu h1 h2. transport
    λw. qSqrtAt u Dm ≤ w + qOfNat (S (isqrt T)) * (qInvNat Dm * qInvNat Dn)
    {qOfNat (S (qSqrtNum v Dn)) * qInvNat Dn}
    {qSqrtAt v Dn + qInvNat Dn}
    qMulSucNat (qSqrtNum v Dn) (qInvNat Dn)
    crossToQ (sqrtCompareN _ nnu h1 h2)