module Generic.Unsafe.PrimAtom.FamilyMacro where

-- Unsafe family derivation specialized to `PrimAtoms`.
--
-- This predates the safe `Generic.Macro.Family` implementation. It generates
-- descriptions and mutually recursive encode/decode helpers for explicit
-- families of datatypes whose non-recursive fields are strings. The roundtrip
-- proofs are declared as postulates, which is why this module lives in the
-- unsafe library. Prefer `Generic.Macro.Family` for safe derivation.

open import Agda.Builtin.Bool using (Bool; true; false)
open import Agda.Primitive using (Set)
open import Agda.Builtin.List using (List; []; _∷_)
open import Agda.Builtin.Nat using (Nat; zero; suc)
open import Agda.Builtin.Reflection as R
  using
    ( Arg
    ; ArgInfo
    ; Clause
    ; Name
    ; Pattern
    ; TC
    ; Telescope
    ; Term
    ; Visibility
    ; abs
    ; arg
    ; arg-info
    ; bindTC
    ; data-type
    ; declareDef
    ; declarePostulate
    ; def
    ; defineFun
    ; freshName
    ; getDefinition
    ; getType
    ; lit
    ; modality
    ; nameErr
    ; pat-lam
    ; pi
    ; primQNameEquality
    ; primShowQName
    ; quantity-ω
    ; reduce
    ; relevant
    ; returnTC
    ; strErr
    ; typeError
    ; unknown
    ; var
    ; visible
    )
open import Agda.Builtin.String using (String; primStringAppend)
open import Agda.Builtin.Unit using (⊤)

open import Generic.Core hiding (natAtom; stringAtom; boolAtom)
open import Generic.PrimAtom

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

_>>=_ : ∀ {A B : Set} → TC A → (A → TC B) → TC B
_>>=_ = bindTC

_>>_ : ∀ {A B : Set} → TC A → TC B → TC B
ma >> mb = ma >>= λ _ → mb

_++_ : String → String → String
_++_ = primStringAppend

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

varg : ∀ {A : Set} → A → Arg A
varg = arg vinfo

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

length : ∀ {A : Set} → List A → Nat
length [] = zero
length (_ ∷ xs) = suc (length xs)

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

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

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

natTerm : Nat → Term
natTerm n = lit (R.nat n)

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

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

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

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

data MaybeNat : Set where
  no  : MaybeNat
  yes : Nat → MaybeNat

data MaybeName : Set where
  noName   : MaybeName
  justName : Name → MaybeName

data FieldInfo : Set where
  atomF : Term → FieldInfo
  recF  : Nat → Term → FieldInfo

fieldType : FieldInfo → Term
fieldType (atomF t) = t
fieldType (recF _ t) = t

record CtorInfo : Set where
  constructor ctor
  field
    ctorName   : Name
    ctorFields : List FieldInfo

open CtorInfo

record TypeInfo : Set where
  constructor tyinfo
  field
    tyName  : Name
    tyIndex : Nat
    tyCtors : List CtorInfo

open TypeInfo

termName : Term → MaybeName
termName (def f []) = justName f
termName (R.con f []) = justName f
termName _ = noName

findNameIndexFrom : Nat → List Name → Name → MaybeNat
findNameIndexFrom i [] fieldName = no
findNameIndexFrom i (ty ∷ tys) fieldName with primQNameEquality fieldName ty
... | true = yes i
... | false = findNameIndexFrom (suc i) tys fieldName

findFamilyIndexFrom : Nat → List Name → Term → MaybeNat
findFamilyIndexFrom i tys fieldTy with termName fieldTy
... | noName = no
... | justName fieldName = findNameIndexFrom i tys fieldName

findFamilyIndex : List Name → Term → MaybeNat
findFamilyIndex tys fieldTy = findFamilyIndexFrom zero tys fieldTy

fieldInfo : List Name → Term → TC FieldInfo
fieldInfo family fieldTy =
  reduce fieldTy >>= λ reduced →
  caseFind (findFamilyIndex family reduced)
  where
  caseFind : MaybeNat → TC FieldInfo
  caseFind no = returnTC (atomF fieldTy)
  caseFind (yes i) = returnTC (recF i fieldTy)

collectFields : Nat → List Name → Name → Name → List FieldInfo → Term → TC CtorInfo
collectFields zero family ty c fields _ =
  typeError (strErr "deriveStringFamily: out of fuel while reading constructor type: " ∷ nameErr c ∷ [])
collectFields (suc fuel) family ty c fields
  (pi (arg (arg-info visible _) fieldTy) (abs _ body)) =
  fieldInfo family fieldTy >>= λ f →
  collectFields fuel family ty c (fields ++ˡ [ f ]) body
collectFields (suc fuel) family ty c fields (pi _ _) =
  typeError (strErr "deriveStringFamily: hidden or instance constructor fields are not supported: " ∷ nameErr c ∷ [])
collectFields (suc fuel) family ty c fields result with termName result
... | justName resultName with primQNameEquality ty resultName
...   | true = returnTC (ctor c fields)
...   | false =
  typeError (strErr "deriveStringFamily: constructor result is not the datatype itself: " ∷ nameErr c ∷ [])
collectFields (suc fuel) family ty c fields result | noName =
  typeError (strErr "deriveStringFamily: constructor result is not the datatype itself: " ∷ nameErr c ∷ [])

macroFuel : Nat
macroFuel = 64

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

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

typeInfo : List Name → Nat → Name → TC TypeInfo
typeInfo family i ty =
  getDefinition ty >>= λ
    { (data-type zero cs) →
        constructorInfos family ty cs >>= λ infos →
        returnTC (tyinfo ty i infos)
    ; (data-type (suc _) _) →
        typeError (strErr "deriveStringFamily: parameterized datatypes are not supported: " ∷ nameErr ty ∷ [])
    ; _ →
        typeError (strErr "deriveStringFamily: expected a datatype: " ∷ nameErr ty ∷ [])
    }

typeInfosFrom : List Name → Nat → List Name → TC (List TypeInfo)
typeInfosFrom family i [] = returnTC []
typeInfosFrom family i (ty ∷ tys) =
  typeInfo family i ty >>= λ info →
  typeInfosFrom family (suc i) tys >>= λ infos →
  returnTC (info ∷ infos)

typeInfos : List Name → TC (List TypeInfo)
typeInfos family = typeInfosFrom family zero family

primAtomsTerm : Term
primAtomsTerm = def (quote PrimAtoms) []

stringAtomTerm : Term
stringAtomTerm = R.con (quote stringAtom) []

fieldDescTerm : FieldInfo → Term
fieldDescTerm (atomF _) =
  R.con (quote fieldAtom) (varg stringAtomTerm ∷ [])
fieldDescTerm (recF i _) =
  R.con (quote fieldRec) (varg (finTerm i) ∷ [])

-- The descriptor generation is the same shape as the safe family macro:
-- each constructor becomes a `con`, string fields become `fieldAtom
-- stringAtom`, and fields whose reduced type is in the explicit family become
-- `fieldRec` with the corresponding family index.
fieldDescTerms : List FieldInfo → List Term
fieldDescTerms [] = []
fieldDescTerms (f ∷ fs) = fieldDescTerm f ∷ fieldDescTerms fs

conDescTerm : CtorInfo → Term
conDescTerm c =
  R.con (quote con) (varg (termList (fieldDescTerms (ctorFields c))) ∷ [])

conDescTerms : List CtorInfo → List Term
conDescTerms [] = []
conDescTerms (c ∷ cs) = conDescTerm c ∷ conDescTerms cs

descClause : TypeInfo → Clause
descClause info =
  R.clause []
    (varg (finPattern (tyIndex info)) ∷ [])
    (R.con (quote dataD) (varg (termList (conDescTerms (tyCtors info))) ∷ []))

descClauses : List TypeInfo → List Clause
descClauses [] = []
descClauses (i ∷ is) = descClause i ∷ descClauses is

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

constructorNamesClause : TypeInfo → Clause
constructorNamesClause info =
  R.clause []
    (varg (finPattern (tyIndex info)) ∷ [])
    (termList (constructorNameTerms (tyCtors info)))

constructorNamesClauses : List TypeInfo → List Clause
constructorNamesClauses [] = []
constructorNamesClauses (i ∷ is) =
  constructorNamesClause i ∷ constructorNamesClauses is

typeNameClause : TypeInfo → Clause
typeNameClause info =
  R.clause []
    (varg (finPattern (tyIndex info)) ∷ [])
    (lit (R.string (primShowQName (tyName info))))

typeNameClauses : List TypeInfo → List Clause
typeNameClauses [] = []
typeNameClauses (i ∷ is) = typeNameClause i ∷ typeNameClauses is

teleFromFields : Nat → List FieldInfo → Telescope
teleFromFields n [] = []
teleFromFields n (f ∷ fs) =
  ("x" , varg (fieldType f)) ∷ teleFromFields (suc n) fs

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

varTermArgs : Nat → List FieldInfo → List (Arg Term)
varTermArgs n [] = []
varTermArgs n (_ ∷ fs) =
  varg (var n []) ∷ varTermArgs (suc n) fs

argsCodeTerm : List Name → Nat → List FieldInfo → Term
argsCodeTerm encodeNames n [] = R.con (quote []ⁱ) []
argsCodeTerm encodeNames n (atomF _ ∷ fs) =
  R.con (quote _∷ⁱ_)
    ( varg (R.con (quote atomIx) (varg (var n []) ∷ []))
    ∷ varg (argsCodeTerm encodeNames (suc n) fs)
    ∷ [] )
argsCodeTerm encodeNames n (recF i _ ∷ fs) =
  R.con (quote _∷ⁱ_)
    ( varg (R.con (quote recIx)
      (varg (def (lookupName encodeNames i) (varg (var n []) ∷ [])) ∷ []))
    ∷ varg (argsCodeTerm encodeNames (suc n) fs)
    ∷ [] )
  where
  lookupName : List Name → Nat → Name
  lookupName (x ∷ xs) zero = x
  lookupName (x ∷ xs) (suc i) = lookupName xs i
  lookupName [] _ = quote String

argsCodePattern : Nat → List FieldInfo → Pattern
argsCodePattern n [] = R.con (quote []ⁱ) []
argsCodePattern n (atomF _ ∷ fs) =
  R.con (quote _∷ⁱ_)
    ( varg (R.con (quote atomIx) (varg (R.var n) ∷ []))
    ∷ varg (argsCodePattern (suc n) fs)
    ∷ [] )
argsCodePattern n (recF _ _ ∷ fs) =
  R.con (quote _∷ⁱ_)
    ( varg (R.con (quote recIx) (varg (R.var n) ∷ []))
    ∷ varg (argsCodePattern (suc n) fs)
    ∷ [] )

decodeArgs : List Name → Nat → List FieldInfo → List (Arg Term)
decodeArgs decodeNames n [] = []
decodeArgs decodeNames n (atomF _ ∷ fs) =
  varg (var n []) ∷ decodeArgs decodeNames (suc n) fs
decodeArgs decodeNames n (recF i _ ∷ fs) =
  varg (def (lookupName decodeNames i) (varg (var n []) ∷ []))
  ∷ decodeArgs decodeNames (suc n) fs
  where
  lookupName : List Name → Nat → Name
  lookupName (x ∷ xs) zero = x
  lookupName (x ∷ xs) (suc i) = lookupName xs i
  lookupName [] _ = quote String

encodeClause : List Name → Nat → CtorInfo → Clause
encodeClause encodeNames n c =
  R.clause
    (teleFromFields zero (ctorFields c))
    (varg (R.con (ctorName c) (varPatternArgs zero (ctorFields c))) ∷ [])
    (R.con (quote nodeIx)
      (varg (listIxTerm n)
      ∷ varg (argsCodeTerm encodeNames zero (ctorFields c))
      ∷ []))

encodeClauses : List Name → Nat → List CtorInfo → List Clause
encodeClauses encodeNames n [] = []
encodeClauses encodeNames n (c ∷ cs) =
  encodeClause encodeNames n c ∷ encodeClauses encodeNames (suc n) cs

decodeClause : List Name → Nat → CtorInfo → Clause
decodeClause decodeNames n c =
  R.clause
    (teleFromFields zero (ctorFields c))
    (varg (R.con (quote nodeIx)
      (varg (listIxPattern n)
      ∷ varg (argsCodePattern zero (ctorFields c))
      ∷ []))
    ∷ [])
    (R.con (ctorName c) (decodeArgs decodeNames zero (ctorFields c)))

decodeClauses : List Name → Nat → List CtorInfo → List Clause
decodeClauses decodeNames n [] = []
decodeClauses decodeNames n (c ∷ cs) =
  decodeClause decodeNames n c ∷ decodeClauses decodeNames (suc n) cs

-- Helper declarations are emitted before definitions so the generated
-- encode/decode functions can call each other recursively. Unlike the safe
-- module, the proof helpers are postulates, so no reflected proof terms are
-- generated here.
finType : Nat → Term
finType n = def (quote Fin) (varg (natTerm n) ∷ [])

tyDescType : Nat → Term
tyDescType n = def (quote TyDesc) (varg primAtomsTerm ∷ varg (natTerm n) ∷ [])

descType : Nat → Term
descType n =
  pi (varg (finType n)) (abs "_" (tyDescType n))

descAt : Name → Nat → Term
descAt descName i = def descName (varg (finTerm i) ∷ [])

codeType : Name → Nat → Nat → Term
codeType descName count i =
  def (quote CodeIx)
    ( varg primAtomsTerm
    ∷ varg (def descName [])
    ∷ varg (descAt descName i)
    ∷ [] )

rootCodeType : Name → Nat → Nat → Term
rootCodeType descName count i =
  def (quote RootCodeIx)
    ( varg primAtomsTerm
    ∷ varg (def descName [])
    ∷ varg (R.con (quote rootData) (varg (finTerm i) ∷ []))
    ∷ [] )

arrow : Term → Term → Term
arrow a b = pi (varg a) (abs "_" b)

eqType : Term → Term → Term
eqType lhs rhs = def (quote _≡_) (varg lhs ∷ varg rhs ∷ [])

decodeEncodeType : Name → Name → Name → TypeInfo → Term
decodeEncodeType descName encodeName decodeName info =
  pi (varg (def (tyName info) []))
    (abs "x"
      (eqType
        (def decodeName (varg (def encodeName (varg (var 0 []) ∷ [])) ∷ []))
        (var 0 [])))

encodeDecodeType : Name → Name → Name → TypeInfo → Term
encodeDecodeType descName encodeName decodeName info =
  pi (varg (codeType descName zero (tyIndex info)))
    (abs "x"
      (eqType
        (def encodeName (varg (def decodeName (varg (var 0 []) ∷ [])) ∷ []))
        (var 0 [])))

declareAll :
  Nat → Name → List TypeInfo → List Name → List Name → List Name → List Name → TC ⊤
declareAll count descName [] [] [] [] [] = returnTC tt
declareAll count descName (info ∷ infos)
  (enc ∷ encs) (dec ∷ decs) (de ∷ des) (ed ∷ eds) =
  declareDef (varg enc) (arrow (def (tyName info) []) (codeType descName count (tyIndex info))) >>
  declareDef (varg dec) (arrow (codeType descName count (tyIndex info)) (def (tyName info) [])) >>
  declarePostulate (varg de) (decodeEncodeType descName enc dec info) >>
  declarePostulate (varg ed) (encodeDecodeType descName enc dec info) >>
  declareAll count descName infos encs decs des eds
declareAll _ _ _ _ _ _ _ =
  typeError (strErr "deriveStringFamily: internal arity mismatch while declaring helpers" ∷ [])

defineAllFrom :
  List Name → List Name → List TypeInfo → List Name → List Name → TC ⊤
defineAllFrom allEncs allDecs [] [] [] = returnTC tt
defineAllFrom allEncs allDecs (info ∷ infos) (enc ∷ encs) (dec ∷ decs) =
  defineFun enc (encodeClauses allEncs zero (tyCtors info)) >>
  defineFun dec (decodeClauses allDecs zero (tyCtors info)) >>
  defineAllFrom allEncs allDecs infos encs decs
defineAllFrom _ _ _ _ _ =
  typeError (strErr "deriveStringFamily: internal arity mismatch while defining helpers" ∷ [])

defineAll :
  List TypeInfo → List Name → List Name → TC ⊤
defineAll infos encs decs = defineAllFrom encs decs infos encs decs

freshNames : String → Nat → TC (List Name)
freshNames prefix zero = returnTC []
freshNames prefix (suc n) =
  freshName (prefix ++ "-helper") >>= λ x →
  freshNames prefix n >>= λ xs →
  returnTC (x ∷ xs)

lookupName : List Name → Nat → Name
lookupName (x ∷ xs) zero = x
lookupName (x ∷ xs) (suc i) = lookupName xs i
lookupName [] _ = quote String

lookupInfo : List TypeInfo → Nat → TypeInfo
lookupInfo (x ∷ xs) zero = x
lookupInfo (x ∷ xs) (suc i) = lookupInfo xs i
lookupInfo [] _ = tyinfo (quote String) zero []

rootEncodeTerm : Name → Term
rootEncodeTerm enc =
  pat-lam
    (R.clause
      (("x" , varg unknown) ∷ [])
      (varg (R.var zero) ∷ [])
      (R.con (quote rootDataIx) (varg (def enc (varg (var zero []) ∷ [])) ∷ []))
    ∷ [])
    []

rootDecodeTerm : Name → Term
rootDecodeTerm dec =
  pat-lam
    (R.clause
      (("x" , varg unknown) ∷ [])
      (varg (R.con (quote rootDataIx) (varg (R.var zero) ∷ [])) ∷ [])
      (def dec (varg (var zero []) ∷ []))
    ∷ [])
    []

rootDecodeEncodeTerm : Name → Term
rootDecodeEncodeTerm de = def de []

rootEncodeDecodeTerm : Name → Term
rootEncodeDecodeTerm ed =
  pat-lam
    (R.clause
      (("x" , varg unknown) ∷ [])
      (varg (R.con (quote rootDataIx) (varg (R.var zero) ∷ [])) ∷ [])
      (R.def (quote cong)
        ( varg (R.con (quote rootDataIx) [])
        ∷ varg (def ed (varg (var zero []) ∷ []))
        ∷ []))
    ∷ [])
    []

genericTerm :
  Nat → Nat → Name → List TypeInfo → List Name → List Name → List Name → List Name → Term
genericTerm count rootIx descName infos encs decs des eds =
  R.con (quote generic)
    ( varg (natTerm count)
    ∷ varg (def descName [])
    ∷ varg (R.con (quote rootData) (varg (finTerm rootIx) ∷ []))
    ∷ varg (pat-lam (typeNameClauses infos) [])
    ∷ varg (pat-lam (constructorNamesClauses infos) [])
    ∷ varg (rootEncodeTerm (lookupName encs rootIx))
    ∷ varg (rootDecodeTerm (lookupName decs rootIx))
    ∷ varg (rootDecodeEncodeTerm (lookupName des rootIx))
    ∷ varg (rootEncodeDecodeTerm (lookupName eds rootIx))
    ∷ [] )

deriveStringFamilyInTC : List Name → Name → Term → TC ⊤
deriveStringFamilyInTC family root hole =
  typeInfos family >>= λ infos →
  freshName "string-family-desc" >>= λ descName →
  freshNames "encode" (length family) >>= λ encs →
  freshNames "decode" (length family) >>= λ decs →
  freshNames "decode-encode" (length family) >>= λ des →
  freshNames "encode-decode" (length family) >>= λ eds →
  let count = length family in
  declareDef (varg descName) (descType count) >>
  defineFun descName (descClauses infos) >>
  declareAll count descName infos encs decs des eds >>
  defineAll infos encs decs >>
  caseRoot (findFamilyIndex family (def root [])) λ rootIx →
  R.unify hole (genericTerm count rootIx descName infos encs decs des eds)
  where
  caseRoot : MaybeNat → (Nat → TC ⊤) → TC ⊤
  caseRoot no k =
    typeError (strErr "deriveStringFamily: root type is not in family: " ∷ nameErr root ∷ [])
  caseRoot (yes i) k = k i

-- Public unsafe entry points for one- and two-type string families. They are
-- intentionally narrower than the safe module and kept for the interactive
-- prim-atom examples that use postulated proof helpers.
macro
  deriveStringFamily1 : Name → Term → TC ⊤
  deriveStringFamily1 ty hole =
    deriveStringFamilyInTC (ty ∷ []) ty hole

  deriveStringFamily2 : Name → Name → Name → Term → TC ⊤
  deriveStringFamily2 a b root hole =
    deriveStringFamilyInTC (a ∷ b ∷ []) root hole