Natural.sqrt

-- THE INTEGER SQUARE ROOT, and the two facts that pin it down:
--
--   isqrt m * isqrt m ≤ m        (it is a lower approximant)
--   m < S (isqrt m) * S (isqrt m)   (it is the LARGEST one)
--
-- Everything ℝ's square root will need bottoms out here. The rational
-- approximation to √u at depth N is isqrt of ⌊u·N²⌋ divided by N, and
-- the two bounds above are exactly what turns into |r² − u| ≤ 2/N.
--
-- `isqrtAux m f` searches downwards from f for the largest k with
-- k² ≤ m, so `isqrt m = isqrtAux m m` is correct because k² ≤ m
-- forces k ≤ m. The test is `S f * S f ∸ m ≡ Z`, decided by natCase:
-- monus is a FUNCTION, so the branch taken is unambiguous even when
-- the two sides are equal — using leTotal here would be wrong, since
-- at k² = m both disjuncts hold and the wrong one loses the answer.
-- ===== strict comparison =====
-- a ≤ b splits into "b ≤ a is an equality" and "b < a", by casing on
-- the witness LeN carries

import Natural (+, *, plusZeroId, plusAssoc, plusComm, sucPlus, plusSucId, multComm, multDistrib, multSucId, sucMult)
import Natural.order (≤, leRefl, leOfEq, leZero, leZeroInv, leTrans, leSucSelf, leSucMono, leSucInv, leTotal, leMultMonoR, lePlusMonoL)
import Natural.more (∸, multDistribR)
import Natural.div (natCase, natCaseElimEqD, leOfMonusZero, ltOfMonusSuc, leSucNotSelf, ltProdSuc)
import Core.equality (trans, sym, cong, transport)
import Core.id (Id, idToEq, eqToId)

leSplit : {b a : ℕ} → b ≤ a → a ≤ b ⊎ S b ≤ a using (Natural.order.≤.unfold)
leSplit =
  λb a h. ℕ-elim
    k. (b + k ≡ a) → a ≤ b ⊎ S b ≤ a
    λe. inj₁ (leOfEq (sym _ _ (trans _ _ _ (sym _ _ (plusZeroId b)) e)))
    j rec. λe. inj₂
      j, eqToId _ _ (trans _ _ _ (trans _ _ _ (sucPlus b j) (sym _ _ (plusSucId b j))) e)
    h .π₁
    idToEq _ _ _ (h .π₂)

natCompare : (a b : ℕ) → a ≤ b ⊎ S b ≤ a
natCompare = λa b. ⊎-elim (h. inj₁ h) (h. leSplit h) (leTotal a b)

-- ===== squares are monotone =====
leMultMono : {a b c d : ℕ} → a ≤ b → c ≤ d → a * c ≤ b * d using (Natural.order.≤.unfold)
leMultMono =
  λa b c d h1 h2. leTrans
    _
    _
    _
    leMultMonoR c h1
    transport
      λw. b * c ≤ w
      multComm b d
      transport (λw. w ≤ d * b) (multComm b c) (leMultMonoR b h2)

leSquare : {a b : ℕ} → a ≤ b → a * a ≤ b * b
leSquare = λa b h. leMultMono h h

-- a positive number is at most its own square, which is why searching
-- up to m is enough
leSelfSquare : (j : ℕ) → S j ≤ S j * S j using (Natural.order.≤.unfold)
leSelfSquare = λj. j * S j, eqToId _ _ (sym _ _ (sucMult j (S j)))

-- ===== the search =====
isqrtAux : ℕ → ℕ → ℕ
isqrtAux = λm f. ℕ-elim Z (f' ih. natCase (S f') (λw. ih) (S f' * S f' ∸ m)) f

isqrt : ℕ → ℕ
isqrt = λm. isqrtAux m m

-- ===== it is a lower approximant =====
isqrtAuxLe : {m f : ℕ} → isqrtAux m f * isqrtAux m f ≤ m using (isqrtAux.eq, Natural.*.eq)
isqrtAuxLe =
  λm f. ℕ-elim
    leZero m
    f' ih. natCaseElimEqD
      _
      λr. r * r ≤ m
      S f'
      λw. isqrtAux m f'
      S f' * S f' ∸ m
      λe. leOfMonusZero e
      λj e. ih
    f

-- ===== ...and the largest one =====
isqrtAuxMax : {m f k : ℕ} → k * k ≤ m → k ≤ f → k ≤ isqrtAux m f using (isqrtAux.eq)
isqrtAuxMax =
  λm f. ℕ-elim
    λk hk hf. hf
    f' ih. λk hk hf. natCaseElimEqD
      _
      λr. k ≤ r
      S f'
      λw. isqrtAux m f'
      S f' * S f' ∸ m
      λe. hf
      λj e. ⊎-elim
        w. k ≤ isqrtAux m f'
        h. ih k hk h
        h. 𝟘-elim
          leSucNotSelf _ (leTrans _ _ _ (ltOfMonusSuc _ e) (leTrans _ _ _ (leSquare h) hk))
        natCompare k f'
    f

isqrtLe : (m : ℕ) → isqrt m * isqrt m ≤ m using (isqrt.eq)
isqrtLe = λm. isqrtAuxLe

isqrtMax : {m k : ℕ} → k * k ≤ m → k ≤ m → k ≤ isqrt m using (isqrt.eq)
isqrtMax = λm k h1 h2. isqrtAuxMax h1 h2

-- the upper bound: if S (isqrt m) squared still fit under m, isqrt m
-- would have found it, and would be at least its own successor
isqrtUpper : (m : ℕ) → S m ≤ S (isqrt m) * S (isqrt m)
isqrtUpper =
  λm. ⊎-elim
    h. 𝟘-elim (leSucNotSelf _ (isqrtMax h (leTrans _ _ _ (leSelfSquare (isqrt m)) h)))
    h. h
    natCompare (S (isqrt m) * S (isqrt m)) m

-- ===== squares reflect the order =====
--
-- This is the fact ℝ's square root actually runs on. Comparing two
-- rational approximants r_m and r_n means comparing K_m/(D_m+1) with
-- (K_n+1)/(D_n+1), and cross-multiplying turns that into a comparison
-- of NATURALS — where "x² ≤ y² implies x ≤ y" is available (below)
-- and its ℚ counterpart is not: ℚ has no strict multiplication
-- monotonicity in the corpus, and getting it would mean sign
-- multiplicativity or division. At ℕ, natCompare settles it.
ltSqSuc : (y : ℕ) → S (y * y) ≤ S y * S y
ltSqSuc = λy. leTrans _ _ _ (leSucMono (leMultMono (leRefl y) (leSucSelf y))) (ltProdSuc y y)

leOfSqLe : {x y : ℕ} → x * x ≤ y * y → x ≤ y
leOfSqLe =
  λx y h. ⊎-elim
    le. le
    lt. 𝟘-elim (leSucNotSelf _ (leTrans _ _ _ (leTrans _ _ _ (ltSqSuc y) (leSquare lt)) h))
    natCompare x y

-- ===== a square that overshoots by T overshoots by isqrt T =====
plus4Swap : {a b p q : ℕ} → a + b + (p + q) ≡ a + p + (q + b)
plus4Swap =
  λa b p q. a + b + (p + q)
    ≡⟨ plusAssoc a b (p + q) ⟩ a + (b + (p + q))
    ≡⟨ cong (λw. ℕ) (λw. a + w) (sym _ _ (plusAssoc b p q)) ⟩ a + (b + p + q)
    ≡⟨ cong (λw. ℕ) (λw. a + (w + q)) (plusComm p b) ⟩ a + (p + b + q)
    ≡⟨ cong (λw. ℕ) (λw. a + w) (plusAssoc p b q) ⟩ a + (p + (b + q))
    ≡⟨ cong (λw. ℕ) (λw. a + (p + w)) (plusComm q b) ⟩ a + (p + (q + b))
    ≡⟨ sym _ _ (plusAssoc a p (q + b)) ⟩ a + p + (q + b)

sqExpand : (y c : ℕ) → (y + c) * (y + c) ≡ y * y + y * c + (y * c + c * c)
sqExpand =
  λy c. (y + c) * (y + c)
    ≡⟨ multDistribR (y + c) y c ⟩ y * (y + c) + c * (y + c)
    ≡⟨ cong (λw. ℕ) (λw. w + c * (y + c)) (multDistrib y y c) ⟩ y * y + y * c + c * (y + c)
    ≡⟨ cong (λw. ℕ) (λw. y * y + y * c + w) (multDistrib c y c) ⟩ y * y + y * c + (c * y + c * c)
    ≡⟨ cong (λw. ℕ) (λw. y * y + y * c + (w + c * c)) (multComm y c) ⟩
      y * y + y * c + (y * c + c * c)

sqSumLe : (y c : ℕ) → y * y + c * c ≤ (y + c) * (y + c) using (Natural.order.≤.unfold)
sqSumLe =
  λy c. (,)
    y * c + y * c
    eqToId _ _ (trans (y * y + c * c + (y * c + y * c)) _ _ plus4Swap (sym _ _ (sqExpand y c)))

sqrtStepN : {x y T : ℕ} → x * x ≤ y * y + T → x ≤ y + S (isqrt T)
sqrtStepN =
  λx y T h. leOfSqLe
    leTrans
      _
      _
      _
      h
      leTrans
        _
        _
        _
        lePlusMonoL (y * y) (leTrans _ _ _ (leSucSelf T) (isqrtUpper T))
        sqSumLe y (S (isqrt T))

-- ===== isqrt is strictly under any c whose square overshoots =====
isqrtLtOfLt : (T : ℕ) {c : ℕ} → S T ≤ c * c → S (isqrt T) ≤ c
isqrtLtOfLt =
  λT c h. ⊎-elim
    le. 𝟘-elim (leSucNotSelf _ (leTrans _ _ _ h (leTrans _ _ _ (leSquare le) (isqrtLe T))))
    lt. lt
    natCompare c (isqrt T)

-- ...and a sum of two POSITIVE squares is strictly under the square of
-- the sum, which is what makes isqrt of the error term fit
plus4Mid : (a p q b : ℕ) → a + p + (q + b) ≡ a + b + (p + q)
plus4Mid =
  λa p q b. a + p + (q + b)
    ≡⟨ plusAssoc a p (q + b) ⟩ a + (p + (q + b))
    ≡⟨ cong (λw. ℕ) (λw. a + w) (sym _ _ (plusAssoc p q b)) ⟩ a + (p + q + b)
    ≡⟨ cong (λw. ℕ) (λw. a + w) (plusComm b (p + q)) ⟩ a + (b + (p + q))
    ≡⟨ sym _ _ (plusAssoc a b (p + q)) ⟩ a + b + (p + q)

sucMultSuc : (x y : ℕ) → S x * S y ≡ S (x + S x * y)
sucMultSuc =
  λx y. S x * S y ≡⟨ multSucId (S x) y ⟩ S x + S x * y ≡⟨ sucPlus x (S x * y) ⟩ S (x + S x * y)

sqSumStrict : (x y : ℕ) → S (S x * S x + S y * S y) ≤ (S x + S y) * (S x + S y)
  using (Natural.order.≤.unfold)
sqSumStrict =
  λx y. (,)
    x + S x * y + S (x + S x * y)
    eqToId
      _
      _
      trans
        _
        _
        _
        trans
          _
          _
          _
          sucPlus (S x * S x + S y * S y) (x + S x * y + S (x + S x * y))
          trans
            _
            _
            _
            sym _ _ (plusSucId (S x * S x + S y * S y) (x + S x * y + S (x + S x * y)))
            cong
              λw. ℕ
              λw. S x * S x + S y * S y + w
              trans
                _
                _
                _
                sym _ _ (sucPlus (x + S x * y) (S (x + S x * y)))
                cong (λw. ℕ) (λw. w + w) (sym _ _ (sucMultSuc x y))
        trans
          _
          _
          _
          sym _ _ (plus4Mid (S x * S x) (S x * S y) (S x * S y) (S y * S y))
          sym _ _ (sqExpand (S x) (S y))