\import Algebra.Meta
\import Algebra.Monoid
\import Algebra.Ring
\import Algebra.Semiring
\import Arith.Nat
\import Combinatorics.Factorial
\import Data.Array
\import Order.PartialOrder
\import Order.StrictOrder
\import Paths
\import Paths.Meta
\open Monoid (pow)
\open AddMonoid (BigSum)

\func binom (n i : Nat) : Nat
  | _, 0 => 1
  | 0, _ => 0
  | suc n, suc i => binom n i + binom n (suc i)
  \where {
    \lemma binom_0 {n : Nat} : binom n 0 = 1 \elim n
      | 0 => idp
      | suc n => idp

    \lemma binom_< {n m : Nat} (p : n < m) : binom n m = 0 \elim n, m, p
      | 0, suc m, _ => idp
      | suc n, suc m, NatOrder.suc<suc n<m => pmap2 (+) (binom_< n<m) (binom_< (n<m <∘ id<suc))

    {- | Diagonal: `binom n n = 1`. Pascal gives `binom (suc n) (suc n) = binom n n + binom n (suc n)`;
         the second summand vanishes by `binom_<` since `n < suc n`. -}
    \lemma binom_n_n {n : Nat} : binom n n = 1 \elim n
      | 0 => idp
      | suc n => rewrite (binom_n_n, binom_< id<suc) idp

    {- | Factorial identity (sum form): `binom (k + m) k * fac k * fac m = fac (k + m)`.
         The `(k, m)` parameterisation avoids truncated subtraction. Proof is by
         simultaneous induction on `(k, m)` using IH at `(k, suc m)` and `(suc k, m)`,
         both structurally smaller than `(suc k, suc m)`. -}
    \lemma fac-id (k m : Nat) : binom (k + m) k * fac k * fac m = fac (k + m) \elim k, m
      | 0, m => rewrite binom_0 NatSemiring.ide-left
      | suc k, 0 => rewrite (binom_n_n {suc k}) NatSemiring.ide-left
      | suc k, suc m =>
          \have
            | ih1 : binom (k + suc m) k * fac (suc k) * fac (suc m) = suc k * fac (k + suc m)
                => equation.cSemiring *> pmap (suc k *) (fac-id k (suc m))
            | ih2 : binom (k + suc m) (suc k) * fac (suc k) * fac (suc m) = suc m * fac (k + suc m)
                => equation.cSemiring *> pmap (suc m *) (fac-id (suc k) m)
          \in pmap (* _) NatSemiring.rdistr *> NatSemiring.rdistr *> pmap2 (+) ih1 ih2 *> inv NatSemiring.rdistr

    {- | Factorial identity (truncated-subtraction form):
         `binom n k * fac k * fac (n -' k) = fac n` whenever `k <= n`.
         Derived from {fac-id} by transporting along `k + (n -' k) = n`. -}
    \lemma fac-id-le {n k : Nat} (p : k <= n) : binom n k * fac k * fac (n -' k) = fac n
      => transport (\lam x => binom x k * fac k * fac (n -' k) = fac x) (<=_exists p) (fac-id k (n -' k))

    \lemma expansion-noncomm {R : Ring} {x y : R} (xy=yx : x * y = y * x) {n : Nat}
      : pow (x + y) n = BigSum (\new Array R (suc n) (\lam i => natCoef (binom n i) * (pow x i * pow y (n -' i)))) \elim n
      | 0 => equation.ring
      | suc n => \let | xy i => pow x (suc i) * pow y (n -' i)
                      | cxy (i : Fin (suc n)) => natCoef (binom n i) * (pow x i * pow y (n -' i))
                      | lem1 i : cxy i * x = natCoef (binom n i) * xy i
                          => *-assoc *> pmap (_ *) (*-assoc *> pmap (pow x i *) (Monoid.pow-comm1 (inv xy=yx)) *> inv *-assoc)
                      | t1 : pow (x + y) n * x = BigSum (\new Array R (suc n) (\lam i => natCoef (binom n i) * xy i))
                          => pmap (* x) (expansion-noncomm xy=yx) *> R.BigSum-rdistr {\new Array R.E (suc n) cxy}
                          *> path (\lam j => BigSum (\new Array R (suc n) (\lam i => lem1 i j)))
                      | lem2 i : cxy i * y = natCoef (binom n i) * (pow x i * pow y (suc (n -' i)))
                          => *-assoc *> pmap (_ *) *-assoc
                      | t2 : pow (x + y) n * y = BigSum (\new Array R (suc (suc n)) (\lam i => natCoef (binom n i) * (pow x i * pow y (suc n -' i))))
                          => pmap (* y) (expansion-noncomm xy=yx) *> R.BigSum-rdistr {\new Array R (suc n) cxy}
                          *> path (\lam j => BigSum \lam i => (lem2 i *> pmap (\lam m => _ * (_ * pow y m)) (inv (-'_suc (<_suc_<= (fin_< i))))) j)
                          *> inv (pmap (_ +) (pmap (* _) (pmap natCoef (binom_< (transportInv (n <) (mod_< id<suc) id<suc)) *> natCoefZero) *> zro_*-left) *> zro-right)
                          *> inv (R.BigSum_suc {suc n} {\lam i => natCoef (binom n i) * (pow x i * pow y (suc n -' i))})
                 \in ldistr *> pmap2 (+) t1 (t2 *> pmap (natCoef __ * _ + _) binom_0) *> +-comm *> +-assoc *> pmap (_ +) +-comm
                      *> inv (pmap (_ +) (R.BigSum_+ {suc n} {mkArray (\lam i => natCoef (binom n i) * xy i)} {mkArray (\lam i => natCoef (binom n (suc i)) * xy i)}))
                      *> inv (path (\lam j => _ + BigSum (\new Array R (suc n) (\lam i => (pmap (* xy i) (R.natCoef_+ {binom n i} {binom n (suc i)}) *> rdistr) j))))

    \lemma expansion {R : CRing} {x y : R} {n : Nat}
      : pow (x + y) n = BigSum (\new Array R (suc n) (\lam i => natCoef (binom n i) * (pow x i * pow y (n -' i))))
      => expansion-noncomm *-comm
  }