module Generic.Macro where

open import Generic.Core
open import Generic.Certified

import Agda.Builtin.Reflection as R
open import Agda.Builtin.Bool using (Bool; true; false)

-- This module implements the small, non-recursive derivation macros.
--
-- The reflected code below constructs a complete `Generic` record directly as
-- a term. The supported fragment is intentionally simple: one datatype or one
-- record constructor, visible fields only, no recursive fields. In this
-- fragment both roundtrip proofs reduce to `refl`, because encode/decode only
-- wrap and unwrap constructor arguments in the same order.

infixl 1 _>>=_
infixl 1 _>>_

_>>=_ : ∀ {ℓ ℓ'} {A : Type ℓ} {B : Type ℓ'} → R.TC A → (A → R.TC B) → R.TC B
_>>=_ = R.bindTC

_>>_ : ∀ {ℓ ℓ'} {A : Type ℓ} {B : Type ℓ'} → R.TC A → R.TC B → R.TC B
f >> g = f >>= λ _ → g

vinfo : R.ArgInfo
vinfo = R.arg-info R.visible (R.modality R.relevant R.quantity-ω)

varg : ∀ {ℓ} {A : Type ℓ} → A → R.Arg A
varg = R.arg vinfo

termList : List R.Term → R.Term
termList [] = R.con (quote []) []
termList (x ∷ xs) = R.con (quote _∷_) (varg x ∷ varg (termList xs) ∷ [])

emptyFields : R.Term
emptyFields = termList []

record CtorInfo : Type₀ where
  constructor ctor
  field
    ctorName   : R.Name
    ctorFields : List R.Term

open CtorInfo

_++_ : ∀ {ℓ} {A : Type ℓ} → List A → List A → List A
[] ++ ys = ys
(x ∷ xs) ++ ys = x ∷ xs ++ ys

[_] : ∀ {ℓ} {A : Type ℓ} → A → List A
[ x ] = x ∷ []

listIxTerm : ℕ → R.Term
listIxTerm zero = R.con (quote here) []
listIxTerm (suc n) = R.con (quote there) (varg (listIxTerm n) ∷ [])

listIxPattern : ℕ → R.Pattern
listIxPattern zero = R.con (quote here) []
listIxPattern (suc n) = R.con (quote there) (varg (listIxPattern n) ∷ [])

finTerm : ℕ → R.Term
finTerm zero = R.con (quote Fin.zero) []
finTerm (suc n) = R.con (quote Fin.suc) (varg (finTerm n) ∷ [])

finPattern : ℕ → R.Pattern
finPattern zero = R.con (quote Fin.zero) []
finPattern (suc n) = R.con (quote Fin.suc) (varg (finPattern n) ∷ [])

rootTerm : R.Term
rootTerm = finTerm zero

unitArgsTerm : R.Term
unitArgsTerm = R.con (quote []ⁱ) []

unitArgsPattern : R.Pattern
unitArgsPattern = R.con (quote []ⁱ) []

length : ∀ {ℓ} {A : Type ℓ} → List A → ℕ
length [] = zero
length (_ ∷ xs) = suc (length xs)

_+_ : ℕ → ℕ → ℕ
zero + m = m
suc n + m = suc (n + m)

fieldAtomTerms : List R.Term → List R.Term
fieldAtomTerms [] = []
fieldAtomTerms (x ∷ xs) = R.con (quote fieldAtom) (varg x ∷ []) ∷ fieldAtomTerms xs

inferredAtomCodes : ℕ → List R.Term → List R.Term
inferredAtomCodes _ [] = []
inferredAtomCodes n (_ ∷ xs) = finTerm n ∷ inferredAtomCodes (suc n) xs

unknownAtomCodes : List R.Term → List R.Term
unknownAtomCodes [] = []
unknownAtomCodes (_ ∷ xs) = R.unknown ∷ unknownAtomCodes xs

-- Constructor descriptions are generated from field types plus a strategy for
-- choosing atom codes. `deriveGeneric` uses `Fin` codes for the inferred atom
-- universe. `deriveGenericIn` emits unknown atom codes and relies on Agda
-- unification against the expected `Generic U A` type to choose the concrete
-- codes in the user supplied universe.
conDescTerm : List R.Term → R.Term
conDescTerm atomCodes =
  R.con (quote con) (varg (termList (fieldAtomTerms atomCodes)) ∷ [])

conDescTermsWith : (ℕ → List R.Term → List R.Term) → ℕ → List CtorInfo → List R.Term
conDescTermsWith _ _ [] = []
conDescTermsWith fieldCodes n (c ∷ cs) =
  conDescTerm (fieldCodes n (ctorFields c))
    ∷ conDescTermsWith fieldCodes (n + length (ctorFields c)) cs

constructorDescsTermWith : (ℕ → List R.Term → List R.Term) → List CtorInfo → R.Term
constructorDescsTermWith fieldCodes cs = termList (conDescTermsWith fieldCodes zero cs)

constructorNameTerms : List CtorInfo → List R.Term
constructorNameTerms [] = []
constructorNameTerms (c ∷ cs) =
  R.lit (R.string (R.primShowQName (ctorName c))) ∷ constructorNameTerms cs

allFields : List CtorInfo → List R.Term
allFields [] = []
allFields (c ∷ cs) = ctorFields c ++ allFields cs

descClausesWith : (ℕ → List R.Term → List R.Term) → List CtorInfo → List R.Clause
descClausesWith fieldCodes cs =
  R.clause [] (varg (finPattern zero) ∷ [])
    (R.con (quote dataD) (varg (constructorDescsTermWith fieldCodes cs) ∷ []))
    ∷ []

descTermWith : (ℕ → List R.Term → List R.Term) → List CtorInfo → R.Term
descTermWith fieldCodes cs = R.pat-lam (descClausesWith fieldCodes cs) []

typeNameTerm : R.Name → R.Term
typeNameTerm ty =
  R.pat-lam
    ( R.clause [] (varg (finPattern zero) ∷ [])
      (R.lit (R.string (R.primShowQName ty)))
    ∷ [] )
    []

constructorNamesTerm : List CtorInfo → R.Term
constructorNamesTerm cs =
  R.pat-lam
    ( R.clause [] (varg (finPattern zero) ∷ [])
      (termList (constructorNameTerms cs))
    ∷ [] )
    []

constructorNamesFitTerm : List CtorInfo → R.Term
constructorNamesFitTerm [] = R.con (quote []ⁿ) []
constructorNamesFitTerm (_ ∷ cs) =
  R.con (quote _∷ⁿ_) (varg (constructorNamesFitTerm cs) ∷ [])

constructorNamesFitFunctionTerm : List CtorInfo → R.Term
constructorNamesFitFunctionTerm cs =
  R.pat-lam
    ( R.clause [] (varg (finPattern zero) ∷ [])
      (constructorNamesFitTerm cs)
    ∷ [] )
    []

atomCodeTerm : List R.Term → R.Term
atomCodeTerm fields = R.def (quote Fin) (varg (R.lit (R.nat (length fields))) ∷ [])

atomClauses : ℕ → List R.Term → List R.Clause
atomClauses _ [] = []
atomClauses n (f ∷ fs) =
  R.clause [] (varg (finPattern n) ∷ []) f
    ∷ atomClauses (suc n) fs

atomTerm : List R.Term → R.Term
atomTerm [] =
  R.pat-lam (R.absurd-clause [] (varg (R.absurd 0) ∷ []) ∷ []) []
atomTerm fields = R.pat-lam (atomClauses zero fields) []

atomsTerm : List CtorInfo → R.Term
atomsTerm cs =
  let fields = allFields cs in
  R.con (quote atomUniverse)
    (varg (atomCodeTerm fields) ∷ varg (atomTerm fields) ∷ [])

-- The generated encode/decode clauses mirror each constructor. The de Bruijn
-- variables introduced in the telescope are reused in the constructor pattern
-- and in the generic-code term, so the order of `teleFromFields`,
-- `varPatternArgs`, `argsCodeTerm`, and `varTermArgs` must stay aligned.
teleFromFields : ℕ → List R.Term → R.Telescope
teleFromFields _ [] = []
teleFromFields n (f ∷ fs) = ("x" , varg f) ∷ teleFromFields (suc n) fs

varPatternArgs : ℕ → List R.Term → List (R.Arg R.Pattern)
varPatternArgs _ [] = []
varPatternArgs n (_ ∷ fs) = varg (R.var n) ∷ varPatternArgs (suc n) fs

varTermArgs : ℕ → List R.Term → List (R.Arg R.Term)
varTermArgs _ [] = []
varTermArgs n (_ ∷ fs) = varg (R.var n []) ∷ varTermArgs (suc n) fs

argsCodeTerm : ℕ → List R.Term → R.Term
argsCodeTerm _ [] = unitArgsTerm
argsCodeTerm n (_ ∷ fs) =
  R.con (quote _∷ⁱ_)
    ( varg (R.con (quote atomIx) (varg (R.var n []) ∷ []))
    ∷ varg (argsCodeTerm (suc n) fs)
    ∷ [] )

argsCodePattern : ℕ → List R.Term → R.Pattern
argsCodePattern _ [] = unitArgsPattern
argsCodePattern n (_ ∷ fs) =
  R.con (quote _∷ⁱ_)
    ( varg (R.con (quote atomIx) (varg (R.var n) ∷ []))
    ∷ varg (argsCodePattern (suc n) fs)
    ∷ [] )

nodeTerm : ℕ → List R.Term → R.Term
nodeTerm n fields =
  R.con (quote rootDataIx)
    (varg (R.con (quote nodeIx) (varg (listIxTerm n) ∷ varg (argsCodeTerm zero fields) ∷ [])) ∷ [])

nodePattern : ℕ → List R.Term → R.Pattern
nodePattern n fields =
  R.con (quote rootDataIx)
    (varg (R.con (quote nodeIx) (varg (listIxPattern n) ∷ varg (argsCodePattern zero fields) ∷ [])) ∷ [])

encodeClauses : ℕ → List CtorInfo → List R.Clause
encodeClauses _ [] = []
encodeClauses n (c ∷ cs) =
  R.clause
    (teleFromFields zero (ctorFields c))
    (varg (R.con (ctorName c) (varPatternArgs zero (ctorFields c))) ∷ [])
    (nodeTerm n (ctorFields c))
    ∷ encodeClauses (suc n) cs

decodeClauses : ℕ → List CtorInfo → List R.Clause
decodeClauses _ [] = []
decodeClauses n (c ∷ cs) =
  R.clause
    (teleFromFields zero (ctorFields c))
    (varg (nodePattern n (ctorFields c)) ∷ [])
    (R.con (ctorName c) (varTermArgs zero (ctorFields c)))
    ∷ decodeClauses (suc n) cs

decodeEncodeClauses : List CtorInfo → List R.Clause
decodeEncodeClauses [] = []
decodeEncodeClauses (c ∷ cs) =
  R.clause
    (teleFromFields zero (ctorFields c))
    (varg (R.con (ctorName c) (varPatternArgs zero (ctorFields c))) ∷ [])
    (R.def (quote refl) [])
    ∷ decodeEncodeClauses cs

encodeDecodeClauses : ℕ → List CtorInfo → List R.Clause
encodeDecodeClauses _ [] = []
encodeDecodeClauses n (c ∷ cs) =
  R.clause
    (teleFromFields zero (ctorFields c))
    (varg (nodePattern n (ctorFields c)) ∷ [])
    (R.def (quote refl) [])
    ∷ encodeDecodeClauses (suc n) cs

isPi : R.Term → Bool
isPi (R.pi _ _) = true
isPi _ = false

isTarget : R.Name → R.Term → Bool
isTarget ty (R.def f []) = R.primQNameEquality ty f
isTarget ty (R.con f []) = R.primQNameEquality ty f
isTarget _ _ = false

ensureNonIndexedType : R.Name → R.TC Unit
ensureNonIndexedType ty =
  R.getType ty >>= λ tyTy →
  R.reduce tyTy >>= λ rtyTy →
  ensureNotPi ty rtyTy
  where
  ensureNotPi : R.Name → R.Term → R.TC Unit
  ensureNotPi ty tyTy with isPi tyTy
  ... | true =
    R.typeError (R.strErr "deriveGeneric: indexed or parameterized types are not supported" ∷ R.nameErr ty ∷ [])
  ... | false =
    R.returnTC tt

ensureNotRecursiveField : R.Name → R.Name → R.Term → R.TC Unit
ensureNotRecursiveField ty c fieldTy with isTarget ty fieldTy
... | true =
  R.typeError
    ( R.strErr "deriveGeneric: recursive constructor fields are not supported yet: "
    ∷ R.nameErr c
    ∷ [] )
... | false = R.returnTC tt

macroFuel : ℕ
macroFuel = 64

-- Constructor types are read as a visible-argument telescope ending in the
-- datatype itself. The fuel is only a guard against malformed or unexpectedly
-- large reflected terms; it is not part of the generated generic description.
collectFields : ℕ → R.Name → R.Name → List R.Term → R.Term → R.TC CtorInfo
collectFields zero ty c fields _ =
  R.typeError
    ( R.strErr "deriveGeneric: out of fuel while reading constructor type: "
    ∷ R.nameErr c
    ∷ [] )
collectFields (suc fuel) ty c fields (R.pi (R.arg (R.arg-info R.visible _) fieldTy) (R.abs _ body)) =
  R.reduce fieldTy >>= λ reducedFieldTy →
  ensureNotRecursiveField ty c reducedFieldTy >>
  collectFields fuel ty c (fields ++ [ fieldTy ]) body
collectFields (suc fuel) ty c fields (R.pi _ _) =
  R.typeError
    ( R.strErr "deriveGeneric: hidden or instance constructor fields are not supported: "
    ∷ R.nameErr c
    ∷ [] )
collectFields (suc fuel) ty c fields result with isTarget ty result
... | true = R.returnTC (ctor c fields)
... | false =
  R.typeError
    ( R.strErr "deriveGeneric: constructor result is not the datatype itself: "
    ∷ R.nameErr c
    ∷ [] )

constructorInfo : R.Name → R.Name → R.TC CtorInfo
constructorInfo ty c =
  R.getType c >>= λ cTy →
  R.reduce cTy >>= λ reducedCTy →
  collectFields macroFuel ty c [] reducedCTy

constructorInfos : R.Name → List R.Name → R.TC (List CtorInfo)
constructorInfos _ [] = R.returnTC []
constructorInfos ty (c ∷ cs) =
  constructorInfo ty c >>= λ i →
  constructorInfos ty cs >>= λ is →
  R.returnTC (i ∷ is)

genericTermWith : R.Name → (ℕ → List R.Term → List R.Term) → List CtorInfo → R.Term
genericTermWith ty fieldCodes cs =
  R.con (quote generic)
    ( varg (R.lit (R.nat 1))
    ∷ varg (descTermWith fieldCodes cs)
    ∷ varg (R.con (quote rootData) (varg rootTerm ∷ []))
    ∷ varg (typeNameTerm ty)
    ∷ varg (constructorNamesTerm cs)
    ∷ varg (R.pat-lam (encodeClauses zero cs) [])
    ∷ varg (R.pat-lam (decodeClauses zero cs) [])
    ∷ varg (R.pat-lam (decodeEncodeClauses cs) [])
    ∷ varg (R.pat-lam (encodeDecodeClauses zero cs) [])
    ∷ [] )

inferredGenericTerm : R.Name → List CtorInfo → R.Term
inferredGenericTerm ty cs =
  R.con (quote _,_)
    ( varg (atomsTerm cs)
    ∷ varg (genericTermWith ty inferredAtomCodes cs)
    ∷ [] )

certifiedGenericTermWith :
  R.Name → (ℕ → List R.Term → List R.Term) → List CtorInfo → R.Term
certifiedGenericTermWith ty fieldCodes cs =
  R.con (quote certifiedGeneric)
    ( varg (genericTermWith ty fieldCodes cs)
    ∷ varg (constructorNamesFitFunctionTerm cs)
    ∷ [] )

inferredCertifiedGenericTerm : R.Name → List CtorInfo → R.Term
inferredCertifiedGenericTerm ty cs =
  R.con (quote _,_)
    ( varg (atomsTerm cs)
    ∷ varg (certifiedGenericTermWith ty inferredAtomCodes cs)
    ∷ [] )

deriveGenericWith : (R.Name → List CtorInfo → R.Term) → R.Name → R.Term → R.TC Unit
deriveGenericWith mkTerm ty hole =
  ensureNonIndexedType ty >>
  R.getDefinition ty >>= λ
    { (R.data-type zero cs) →
        constructorInfos ty cs >>= λ infos →
        R.noConstraints (R.unify hole (mkTerm ty infos))
    ; (R.record-type c _) →
        constructorInfo ty c >>= λ info →
        R.noConstraints (R.unify hole (mkTerm ty (info ∷ [])))
    ; (R.data-type (suc _) _) →
        R.typeError (R.strErr "deriveGeneric: parameterized datatypes are not supported" ∷ R.nameErr ty ∷ [])
    ; _ →
        R.typeError (R.strErr "deriveGeneric: expected a datatype or record type" ∷ R.nameErr ty ∷ [])
    }

deriveGenericTC : R.Name → R.Term → R.TC Unit
deriveGenericTC = deriveGenericWith inferredGenericTerm

deriveGenericInTC : R.Term → R.Name → R.Term → R.TC Unit
deriveGenericInTC atoms =
  deriveGenericWith
    (λ ty infos →
      R.def (quote withAtomUniverse)
        (varg atoms
        ∷ varg (genericTermWith ty (λ _ fields → unknownAtomCodes fields) infos)
        ∷ []))

deriveCertifiedGenericTC : R.Name → R.Term → R.TC Unit
deriveCertifiedGenericTC = deriveGenericWith inferredCertifiedGenericTerm

deriveCertifiedGenericInTC : R.Term → R.Name → R.Term → R.TC Unit
deriveCertifiedGenericInTC atoms =
  deriveGenericWith
    (λ ty infos →
      R.con (quote certifiedGeneric)
        ( varg
            (R.def (quote withAtomUniverse)
              ( varg atoms
              ∷ varg (genericTermWith ty (λ _ fields → unknownAtomCodes fields) infos)
              ∷ [] ))
        ∷ varg (constructorNamesFitFunctionTerm infos)
        ∷ [] ))

-- Public macro entry points. The last reflected argument is the hole supplied
-- by Agda at the call site; each implementation unifies that hole with the
-- generated term.
macro
  deriveGeneric : R.Name → R.Term → R.TC Unit
  deriveGeneric = deriveGenericTC

  deriveGenericEnum : R.Name → R.Term → R.TC Unit
  deriveGenericEnum = deriveGenericTC

  deriveGenericIn : R.Term → R.Name → R.Term → R.TC Unit
  deriveGenericIn = deriveGenericInTC

  deriveCertifiedGeneric : R.Name → R.Term → R.TC Unit
  deriveCertifiedGeneric = deriveCertifiedGenericTC

  deriveCertifiedGenericEnum : R.Name → R.Term → R.TC Unit
  deriveCertifiedGenericEnum = deriveCertifiedGenericTC

  deriveCertifiedGenericIn : R.Term → R.Name → R.Term → R.TC Unit
  deriveCertifiedGenericIn = deriveCertifiedGenericInTC