module Generic.Unsafe.PrimAtom.FamilyMacro where
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) ∷ [])
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
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
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