open import 1Lab.Reflection

open import Cat.Prelude

import Cat.Reasoning as Cr
import Cat.Solver as Cs

module Cat.Functor.Solver where

open Functor

module NbE where
  open Cs.NbE using (`id ; _โ†‘ ; _`โˆ˜_)

  private
    module CE = Cs.NbE
    variable
      o o' h h' : Level
      ๐’Ÿ : Precategory o h

  data FExpr : (๐’Ÿ : Precategory o h) โ†’ โŒž ๐’Ÿ โŒŸ โ†’ โŒž ๐’Ÿ โŒŸ โ†’ Typeฯ‰ where
    `Fโ‚
      : (๐’ž : Precategory o h) (F : Functor ๐’ž ๐’Ÿ) {A B : โŒž ๐’ž โŒŸ}
      โ†’ FExpr ๐’ž A B โ†’ FExpr ๐’Ÿ (F .Fโ‚€ A) (F .Fโ‚€ B)
    _`โˆ˜_ : {X Y Z : โŒž ๐’Ÿ โŒŸ} โ†’ FExpr ๐’Ÿ Y Z โ†’ FExpr ๐’Ÿ X Y โ†’ FExpr ๐’Ÿ X Z
    `id  : {X : โŒž ๐’Ÿ โŒŸ} โ†’ FExpr ๐’Ÿ X X
    _โ†‘   : {X Y : โŒž ๐’Ÿ โŒŸ} โ†’ Cr.Hom ๐’Ÿ X Y โ†’ FExpr ๐’Ÿ X Y

  unfexpr : (๐’Ÿ : Precategory o h) {X Y : โŒž ๐’Ÿ โŒŸ} โ†’ FExpr ๐’Ÿ X Y โ†’ Cr.Hom ๐’Ÿ X Y
  unfexpr ๐’Ÿ (`Fโ‚ ๐’ž F e) = F .Fโ‚ (unfexpr ๐’ž e)
  unfexpr ๐’Ÿ (e1 `โˆ˜ e2)  = unfexpr ๐’Ÿ e1 โˆ˜ unfexpr ๐’Ÿ e2 where open Precategory ๐’Ÿ
  unfexpr ๐’Ÿ `id         = Cr.id ๐’Ÿ
  unfexpr ๐’Ÿ (f โ†‘)       = f

  --------------------------------------------------------------------------------
  -- Evaluation

  CExpr : (๐’ž : Precategory o h) โ†’ โŒž ๐’ž โŒŸ โ†’ โŒž ๐’ž โŒŸ โ†’ Type (o โŠ” h)
  CExpr = CE.Expr

  do-fmap
    : (๐’ž : Precategory o h) (๐’Ÿ : Precategory o' h') (F : Functor ๐’ž ๐’Ÿ)
    โ†’ {A B : โŒž ๐’ž โŒŸ} โ†’ CExpr ๐’ž A B โ†’ CExpr ๐’Ÿ (F .Fโ‚€ A) (F .Fโ‚€ B)
  do-fmap ๐’ž ๐’Ÿ F `id       = `id
  do-fmap ๐’ž ๐’Ÿ F (e `โˆ˜ eโ‚) = do-fmap ๐’ž ๐’Ÿ F e `โˆ˜ do-fmap ๐’ž ๐’Ÿ F eโ‚
  do-fmap ๐’ž ๐’Ÿ F (f โ†‘)     = F .Fโ‚ f โ†‘

  eval : (๐’Ÿ : Precategory o h) {X Y : โŒž ๐’Ÿ โŒŸ} โ†’ FExpr ๐’Ÿ X Y โ†’ CExpr ๐’Ÿ X Y
  eval ๐’Ÿ (`Fโ‚ ๐’ž F e) = do-fmap ๐’ž ๐’Ÿ F (eval ๐’ž e)
  eval ๐’Ÿ (e1 `โˆ˜ e2)  = eval ๐’Ÿ e1 `โˆ˜ eval ๐’Ÿ e2
  eval ๐’Ÿ `id         = `id
  eval ๐’Ÿ (f โ†‘)       = f โ†‘

  nf : (๐’Ÿ : Precategory o h) {X Y : โŒž ๐’Ÿ โŒŸ} โ†’ FExpr ๐’Ÿ X Y โ†’ Cr.Hom ๐’Ÿ X Y
  nf ๐’Ÿ e = CE.nf ๐’Ÿ (eval ๐’Ÿ e)

  --------------------------------------------------------------------------------
  -- Soundness

  do-fmap-sound
    : (๐’ž : Precategory o h) (๐’Ÿ : Precategory o' h') (F : Functor ๐’ž ๐’Ÿ) {A B : โŒž ๐’ž โŒŸ}
    โ†’ (v : CE.Expr ๐’ž A B) โ†’ CE.embed ๐’Ÿ (do-fmap ๐’ž ๐’Ÿ F v) โ‰ก F .Fโ‚ (CE.embed ๐’ž v)
  do-fmap-sound ๐’ž ๐’Ÿ F `id       = sym (F .F-id)
  do-fmap-sound ๐’ž ๐’Ÿ F (v `โˆ˜ vโ‚) =
    CE.embed ๐’Ÿ (do-fmap ๐’ž ๐’Ÿ F v) ๐’Ÿ.โˆ˜ CE.embed ๐’Ÿ (do-fmap ๐’ž ๐’Ÿ F vโ‚) โ‰กโŸจ apโ‚‚ ๐’Ÿ._โˆ˜_ (do-fmap-sound ๐’ž ๐’Ÿ F v) (do-fmap-sound ๐’ž ๐’Ÿ F vโ‚) โŸฉ
    F .Fโ‚ (CE.embed ๐’ž v) ๐’Ÿ.โˆ˜ F .Fโ‚ (CE.embed ๐’ž vโ‚)                    โ‰กห˜โŸจ F .F-โˆ˜ _ _ โŸฉ
    F .Fโ‚ (CE.embed ๐’ž v ๐’ž.โˆ˜ CE.embed ๐’ž vโ‚)                            โˆŽ
    where
      module ๐’Ÿ = Precategory ๐’Ÿ
      module ๐’ž = Precategory ๐’ž
  do-fmap-sound ๐’ž ๐’Ÿ F (x โ†‘) = refl

  eval-sound
    : (๐’Ÿ : Precategory o h) {X Y : โŒž ๐’Ÿ โŒŸ} โ†’ (e : FExpr ๐’Ÿ X Y)
    โ†’ CE.embed ๐’Ÿ (eval ๐’Ÿ e) โ‰ก unfexpr ๐’Ÿ e
  eval-sound ๐’Ÿ (`Fโ‚ ๐’ž F v) =
    do-fmap-sound ๐’ž ๐’Ÿ F (eval ๐’ž v) โˆ™ ap (F .Fโ‚) (eval-sound ๐’ž v)
  eval-sound ๐’Ÿ (e `โˆ˜ eโ‚) = apโ‚‚ _โˆ˜_ (eval-sound ๐’Ÿ e) (eval-sound ๐’Ÿ eโ‚)
    where open Precategory ๐’Ÿ
  eval-sound ๐’Ÿ `id   = refl
  eval-sound ๐’Ÿ (f โ†‘) = refl

  nf-sound
    : (๐’Ÿ : Precategory o h) {X Y : โŒž ๐’Ÿ โŒŸ} (e : FExpr ๐’Ÿ X Y) โ†’ nf ๐’Ÿ e โ‰ก unfexpr ๐’Ÿ e
  nf-sound ๐’Ÿ e = CE.eval-sound ๐’Ÿ (eval ๐’Ÿ e) โˆ™ eval-sound ๐’Ÿ e

  abstract
    solve
      : (๐’Ÿ : Precategory o h) {X Y : โŒž ๐’Ÿ โŒŸ} โ†’ (e1 e2 : FExpr ๐’Ÿ X Y)
      โ†’ nf ๐’Ÿ e1 โ‰ก nf ๐’Ÿ e2 โ†’ unfexpr ๐’Ÿ e1 โ‰ก unfexpr ๐’Ÿ e2
    solve ๐’Ÿ e1 e2 p = sym (nf-sound ๐’Ÿ e1) โˆ™โˆ™ p โˆ™โˆ™ (nf-sound ๐’Ÿ e2)

module Reflection where

  open Cs.Reflection using (โ€œidโ€ ; โ€œโˆ˜โ€)

  pattern functor-args cat functor xs =
    _ hโˆท _ hโˆท cat hโˆท _ hโˆท _ hโˆท _ hโˆท functor vโˆท xs

  pattern โ€œFโ‚โ€ cat functor f =
    def (quote Functor.Fโ‚) (functor-args cat functor (_ hโˆท _ hโˆท f vโˆท []))

  โ€œsolveโ€ : Term โ†’ Term โ†’ Term โ†’ Term
  โ€œsolveโ€ cat lhs rhs =
    def (quote NbE.solve) (cat vโˆท lhs vโˆท rhs vโˆท def (quote refl) [] vโˆท [])

  build-fexpr : Term โ†’ Term
  build-fexpr โ€œidโ€      = con (quote NbE.FExpr.`id) []
  build-fexpr (โ€œโˆ˜โ€ f g) = con (quote NbE.FExpr._`โˆ˜_)
    (build-fexpr f vโˆท build-fexpr g vโˆท [])
  build-fexpr (โ€œFโ‚โ€ cat functor f) = con (quote NbE.FExpr.`Fโ‚)
    (cat vโˆท functor vโˆท build-fexpr f vโˆท [])
  build-fexpr f = con (quote NbE.FExpr._โ†‘) (f vโˆท [])

  dont-reduce : List Name
  dont-reduce = quote Precategory.id โˆท quote Precategory._โˆ˜_ โˆท quote Functor.Fโ‚ โˆท []

module _ {o h} (๐’Ÿ : Precategory o h) {x y : โŒž ๐’Ÿ โŒŸ} {h1 h2 : ๐’Ÿ .Precategory.Hom x y} where
  open Reflection
  functor-worker : Term โ†’ TC โŠค
  functor-worker hole =
   withNormalisation true $
   withReduceDefs (false , dont-reduce) $ do
     `h1 โ† wait-for-type =<< quoteTC h1
     `h2 โ† quoteTC h2
     `๐’Ÿ โ† quoteTC ๐’Ÿ
     let
       elhs = build-fexpr `h1
       erhs = build-fexpr `h2
     noConstraints $ unify hole (โ€œsolveโ€ `๐’Ÿ elhs erhs)

  functor-wrapper : {@(tactic functor-worker) p : h1 โ‰ก h2} โ†’ h1 โ‰ก h2
  functor-wrapper {p = p} = p

macro
  functor! : Term โ†’ Term โ†’ TC โŠค
  functor! cat = flip unify (def (quote functor-wrapper) (cat vโˆท []))

private
  module Test
    {o h} {๐’ž ๐’Ÿ โ„ฐ : Precategory o h} (F : Functor ๐’ž ๐’Ÿ) (G : Functor ๐’Ÿ โ„ฐ) where
    module ๐’ž = Precategory ๐’ž
    module ๐’Ÿ = Precategory ๐’Ÿ
    module โ„ฐ = Precategory โ„ฐ

    variable
      A B : ๐’ž.Ob
      a b : ๐’ž.Hom A B

    test
      : G .Fโ‚ (F .Fโ‚ a ๐’Ÿ.โˆ˜ ๐’Ÿ.id) โ„ฐ.โˆ˜ G .Fโ‚ (F .Fโ‚ (b ๐’ž.โˆ˜ ๐’ž.id)) โ„ฐ.โˆ˜ โ„ฐ.id
      โ‰ก G .Fโ‚ (F .Fโ‚ (a ๐’ž.โˆ˜ b))
    test = functor! โ„ฐ