Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
2 changes: 1 addition & 1 deletion .gitignore
Original file line number Diff line number Diff line change
Expand Up @@ -23,5 +23,5 @@ dist-newstyle
*.hi
*.chi
*.chs.h

.stack-work

4 changes: 4 additions & 0 deletions dependent-sum-template/ChangeLog.md
Original file line number Diff line number Diff line change
@@ -1,5 +1,9 @@
# Revision history for dependent-sum-template

## 0.2.0.0 - 2026-xx-xx

Upgrade to `some-1.1` and drop support for `GHC < 9.12`.

## 0.1.1.1 - 2021-12-30

* Fix warning with GHC 9.2 about non-canonical `return`.
Expand Down
21 changes: 6 additions & 15 deletions dependent-sum-template/dependent-sum-template.cabal
Original file line number Diff line number Diff line change
@@ -1,5 +1,5 @@
name: dependent-sum-template
version: 0.1.1.1
version: 0.2.0.0
stability: experimental

cabal-version: >= 1.10
Expand All @@ -14,12 +14,7 @@ category: Unclassified
synopsis: Template Haskell code to generate instances of classes in dependent-sum package
description: Template Haskell code to generate instances of classes in dependent-sum package, such as 'GEq' and 'GCompare'.

tested-with: GHC == 8.0.2,
GHC == 8.2.2,
GHC == 8.4.4,
GHC == 8.6.5,
GHC == 8.8.3,
GHC == 9.0.1
tested-with: GHC == 9.12.4

extra-source-files: ChangeLog.md

Expand All @@ -28,22 +23,18 @@ source-repository head
location: https://github.com/obsidiansystems/dependent-sum

Library
if impl(ghc < 7.10)
buildable: False
hs-source-dirs: src
default-language: Haskell2010
exposed-modules: Data.GADT.Compare.TH
Data.GADT.Show.TH
other-modules: Data.Dependent.Sum.TH.Internal
build-depends: base >= 3 && <5,
dependent-sum >= 0.4.1 && < 0.8,
build-depends: base >= 4.21 && <5,
dependent-sum >=0.4.1 && <0.9,
template-haskell,
th-extras >= 0.0.0.2,
th-abstraction >= 0.4
th-extras >=0.0.0.9 && < 0.1,
th-abstraction >=0.7.2.0 && <0.8

test-suite test
if impl(ghc < 8.0)
buildable: False
type: exitcode-stdio-1.0
hs-source-dirs: test
default-language: Haskell2010
Expand Down
10 changes: 5 additions & 5 deletions dependent-sum-template/src/Data/Dependent/Sum/TH/Internal.hs
Original file line number Diff line number Diff line change
Expand Up @@ -26,9 +26,9 @@ classHeadToParams t = (h, reverse reversedParams)
-- we're deriving for is always the first typeclass parameter, if there are
-- multiple.
deriveForDec :: Name -> (Q Type -> Q Type) -> ([TyVarBndrSpec] -> [Con] -> Q Dec) -> Dec -> Q [Dec]
deriveForDec className makeClassHead f dec = deriveForDec' className makeClassHead (f . changeTVFlags specifiedSpec) dec
deriveForDec = deriveForDec'

deriveForDec' :: Name -> (Q Type -> Q Type) -> ([TyVarBndrUnit] -> [Con] -> Q Dec) -> Dec -> Q [Dec]
deriveForDec' :: Name -> (Q Type -> Q Type) -> ([TyVarBndrSpec] -> [Con] -> Q Dec) -> Dec -> Q [Dec]
deriveForDec' className _ f (InstanceD overlaps cxt classHead decs) = do
let (givenClassName, firstParam : _) = classHeadToParams classHead
when (givenClassName /= className) $
Expand All @@ -37,13 +37,13 @@ deriveForDec' className _ f (InstanceD overlaps cxt classHead decs) = do
dataTypeInfo <- reify dataTypeName
case dataTypeInfo of
TyConI (DataD dataCxt name bndrs _ cons _) -> do
dec <- f bndrs cons
dec <- f (changeTVFlags specifiedSpec bndrs) cons
return [InstanceD overlaps cxt classHead [dec]]
_ -> fail $ "while deriving " ++ show className ++ ": the name of an algebraic data type constructor is required"
deriveForDec' className makeClassHead f (DataD dataCxt name bndrs _ cons _) = return <$> inst
where
inst = instanceD (cxt (map return dataCxt)) (makeClassHead $ conT name) [dec]
dec = f bndrs cons
dec = f (changeTVFlags specifiedSpec bndrs) cons
#if __GLASGOW_HASKELL__ >= 808
deriveForDec' className makeClassHead f (DataInstD dataCxt tvBndrs ty _ cons _) = return <$> inst
#else
Expand All @@ -64,4 +64,4 @@ deriveForDec' className makeClassHead f (DataInstD dataCxt name tyArgs _ cons _)
-- TODO: figure out proper number of family parameters vs instance parameters
bndrs = [PlainTV v | VarT v <- tail tyArgs ]
#endif
dec = f bndrs cons
dec = f (changeTVFlags specifiedSpec bndrs) cons
231 changes: 229 additions & 2 deletions dependent-sum-template/src/Data/GADT/Compare/TH.hs
Original file line number Diff line number Diff line change
Expand Up @@ -7,6 +7,8 @@
module Data.GADT.Compare.TH
( DeriveGEQ(..)
, DeriveGCompare(..)
, deriveGEqSuperclasses
, deriveGCompareSuperclasses
, GComparing, runGComparing, geq', compare'
) where

Expand All @@ -20,6 +22,7 @@ import Data.Traversable (for)
import Data.Type.Equality ((:~:) (..))
import Language.Haskell.TH
import Language.Haskell.TH.Extras
import Language.Haskell.TH.Datatype.TyVarBndr

-- A type class purely for overloading purposes
class DeriveGEQ t where
Expand All @@ -33,7 +36,15 @@ instance DeriveGEQ Name where
_ -> fail "deriveGEq: the name of a type constructor is required"

instance DeriveGEQ Dec where
deriveGEq = deriveForDec ''GEq (\t -> [t| GEq $t |]) geqFunction
deriveGEq dec = do
eqPName <- superclassNamed ''GEq "EqP"
eqpName <- classMethodNamed eqPName "eqp"
eqDecs <- deriveEqForDec eqpName dec
liftA2 (++)
(return eqDecs)
(liftA2 (++)
(deriveEqPForDec eqPName eqpName dec)
(deriveForDec ''GEq (\t -> [t| GEq $t |]) geqFunction dec))

instance DeriveGEQ t => DeriveGEQ [t] where
deriveGEq [it] = deriveGEq it
Expand All @@ -42,6 +53,119 @@ instance DeriveGEQ t => DeriveGEQ [t] where
instance DeriveGEQ t => DeriveGEQ (Q t) where
deriveGEq = (>>= deriveGEq)

deriveGEqSuperclasses :: Name -> Q [Dec]
deriveGEqSuperclasses typeName = do
eqPName <- superclassNamed ''GEq "EqP"
eqpName <- classMethodNamed eqPName "eqp"
TyConI (DataD dataCxt name bndrs _ _ _) <- reify typeName
eqType <- applyType name bndrs
eqExists <- isInstance' ''Eq [eqType]
let eqDec = instanceD (cxt (map return dataCxt)) (appT (conT ''Eq) (return eqType))
[geqBoolFunction '(==)]
eqPDec <- instanceD (cxt (map return dataCxt)) (appT (conT eqPName) (conT name))
[geqBoolFunction eqpName]
if eqExists then return [eqPDec] else (: [eqPDec]) <$> eqDec

geqBoolFunction funName = funD funName
[clause [varP x, varP y]
(normalB [| case geq $(varE x) $(varE y) of
Just Refl -> True
Nothing -> False
|]) []]
where
x = mkName "x"
y = mkName "y"

superclassNamed className superclassBase = do
ClassI (ClassD superCxt _ _ _ _) _ <- reify className
case [ name | predType <- superCxt
, let (name, _) = classHeadToParams predType
, nameBase name == superclassBase
] of
[name] -> return name
_ -> fail $ "deriveGEq: could not find superclass " ++ superclassBase ++ " of " ++ show className

classMethodNamed className methodBase = do
ClassI (ClassD _ _ _ _ decs) _ <- reify className
case [ name | SigD name _ <- decs, nameBase name == methodBase ] of
[name] -> return name
_ -> fail $ "deriveGEq: could not find method " ++ methodBase ++ " of " ++ show className

deriveEqForDec eqpName (InstanceD _ cxt classHead _) = do
let (_, firstParam : _) = classHeadToParams classHead
(dataTypeName, fixedArgs) = classHeadToParams firstParam
dataTypeInfo <- reify dataTypeName
case dataTypeInfo of
TyConI (DataD dataCxt name bndrs _ cons _) -> deriveEqInstanceForType eqpName (cxt ++ dataCxt) (instanceType name fixedArgs bndrs) bndrs cons
_ -> return []
deriveEqForDec eqpName (DataD dataCxt name bndrs _ cons _) = deriveEqInstance eqpName dataCxt name bndrs cons
deriveEqForDec _ _ = return []

deriveEqPForDec eqPName eqpName (InstanceD _ instCxt classHead _) = do
let (_, firstParam : _) = classHeadToParams classHead
dataTypeName = headOfType firstParam
dataTypeInfo <- reify dataTypeName
case dataTypeInfo of
TyConI (DataD dataCxt _ bndrs _ cons _) -> (:[]) <$> instanceD (cxt (map return (instCxt ++ dataCxt))) (appT (conT eqPName) (return firstParam))
[eqpFunction eqpName (changeTVFlags specifiedSpec bndrs) cons]
_ -> return []
deriveEqPForDec eqPName eqpName (DataD dataCxt name bndrs _ cons _) = (:[]) <$> instanceD (cxt (map return dataCxt)) (appT (conT eqPName) (conT name))
[eqpFunction eqpName (changeTVFlags specifiedSpec bndrs) cons]
deriveEqPForDec _ _ _ = return []

deriveEqInstance eqpName dataCxt name bndrs cons = do
eqType <- applyType name bndrs
deriveEqInstanceForType eqpName dataCxt (return eqType) bndrs cons

deriveEqInstanceForType eqpName dataCxt eqTypeQ bndrs cons = do
eqType <- eqTypeQ
exists <- isInstance' ''Eq [eqType]
if exists
then return []
else (:[]) <$> instanceD (cxt (map return dataCxt)) (appT (conT ''Eq) (return eqType))
[eqFunction '(==) eqpName (changeTVFlags specifiedSpec bndrs) cons]

instanceType name fixedArgs bndrs = foldl appT (conT name) (map return fixedArgs ++ map (varT . nameOfBinder) (drop (length fixedArgs) bndrs))

applyType name bndrs = foldl appT (conT name) (map (varT . nameOfBinder) bndrs)

isInstance' className types = recover (return False) (isInstance className types)

eqFunction eqName eqpName bndrs cons = funD eqName
( map (eqpClause eqpName bndrs) cons
++ [ clause [wildP, wildP] (normalB [| False |]) []
| length cons /= 1
]
)

eqpFunction eqpName bndrs cons = funD eqpName
( map (eqpClause eqpName bndrs) cons
++ [ clause [wildP, wildP] (normalB [| False |]) []
| length cons /= 1
]
)

eqpClause eqpName bndrs con = do
let argTypes = argTypesOfCon con
needsEqP argType = any ((`occursInType` argType) . nameOfBinder) (bndrs ++ varsBoundInCon con)

nArgs = length argTypes
lArgNames <- replicateM nArgs (newName "x")
rArgNames <- replicateM nArgs (newName "y")

clause [ conP conName (map varP lArgNames)
, conP conName (map varP rArgNames)
]
( normalB $ foldr (\(lArg, rArg, argType) rest ->
[| $(if needsEqP argType
then [| $(varE eqpName) $(varE lArg) $(varE rArg) |]
else [| $(varE lArg) == $(varE rArg) |])
&& $rest |])
[| True |]
(zip3 lArgNames rArgNames argTypes)
) []
where conName = nameOfCon con

geqFunction bndrs cons = funD 'geq
( map (geqClause bndrs) cons
++ [ clause [wildP, wildP] (normalB [| Nothing |]) []
Expand Down Expand Up @@ -107,7 +231,15 @@ instance DeriveGCompare Name where
_ -> fail "deriveGCompare: the name of a type constructor is required"

instance DeriveGCompare Dec where
deriveGCompare = deriveForDec ''GCompare (\t -> [t| GCompare $t |]) gcompareFunction
deriveGCompare dec = do
ordPName <- superclassNamed ''GCompare "OrdP"
comparepName <- classMethodNamed ordPName "comparep"
ordDecs <- deriveOrdForDec comparepName dec
liftA2 (++)
(return ordDecs)
(liftA2 (++)
(deriveOrdPForDec ordPName comparepName dec)
(deriveForDec ''GCompare (\t -> [t| GCompare $t |]) gcompareFunction dec))

instance DeriveGCompare t => DeriveGCompare [t] where
deriveGCompare [it] = deriveGCompare it
Expand All @@ -116,6 +248,101 @@ instance DeriveGCompare t => DeriveGCompare [t] where
instance DeriveGCompare t => DeriveGCompare (Q t) where
deriveGCompare = (>>= deriveGCompare)

deriveGCompareSuperclasses :: Name -> Q [Dec]
deriveGCompareSuperclasses typeName = do
ordPName <- superclassNamed ''GCompare "OrdP"
comparepName <- classMethodNamed ordPName "comparep"
TyConI (DataD dataCxt name bndrs _ _ _) <- reify typeName
ordType <- applyType name bndrs
ordExists <- isInstance' ''Ord [ordType]
let ordDec = instanceD (cxt (map return dataCxt)) (appT (conT ''Ord) (return ordType))
[gcompareOrderingFunction 'compare]
ordPDec <- instanceD (cxt (map return dataCxt)) (appT (conT ordPName) (conT name))
[gcompareOrderingFunction comparepName]
if ordExists then return [ordPDec] else (: [ordPDec]) <$> ordDec

gcompareOrderingFunction funName = funD funName
[clause [varP x, varP y]
(normalB [| case gcompare $(varE x) $(varE y) of
GLT -> LT
GEQ -> EQ
GGT -> GT
|]) []]
where
x = mkName "x"
y = mkName "y"

deriveOrdForDec comparepName (InstanceD _ cxt classHead _) = do
let (_, firstParam : _) = classHeadToParams classHead
(dataTypeName, fixedArgs) = classHeadToParams firstParam
dataTypeInfo <- reify dataTypeName
case dataTypeInfo of
TyConI (DataD dataCxt name bndrs _ cons _) -> deriveOrdInstanceForType comparepName (cxt ++ dataCxt) (instanceType name fixedArgs bndrs) bndrs cons
_ -> return []
deriveOrdForDec comparepName (DataD dataCxt name bndrs _ cons _) = deriveOrdInstance comparepName dataCxt name bndrs cons
deriveOrdForDec _ _ = return []

deriveOrdPForDec ordPName comparepName (InstanceD _ instCxt classHead _) = do
let (_, firstParam : _) = classHeadToParams classHead
dataTypeName = headOfType firstParam
dataTypeInfo <- reify dataTypeName
case dataTypeInfo of
TyConI (DataD dataCxt _ bndrs _ cons _) -> (:[]) <$> instanceD (cxt (map return (instCxt ++ dataCxt))) (appT (conT ordPName) (return firstParam))
[ordpFunction comparepName (changeTVFlags specifiedSpec bndrs) cons]
_ -> return []
deriveOrdPForDec ordPName comparepName (DataD dataCxt name bndrs _ cons _) = (:[]) <$> instanceD (cxt (map return dataCxt)) (appT (conT ordPName) (conT name))
[ordpFunction comparepName (changeTVFlags specifiedSpec bndrs) cons]
deriveOrdPForDec _ _ _ = return []

deriveOrdInstance comparepName dataCxt name bndrs cons = do
ordType <- applyType name bndrs
deriveOrdInstanceForType comparepName dataCxt (return ordType) bndrs cons

deriveOrdInstanceForType comparepName dataCxt ordTypeQ bndrs cons = do
ordType <- ordTypeQ
exists <- isInstance' ''Ord [ordType]
if exists
then return []
else (:[]) <$> instanceD (cxt (map return dataCxt)) (appT (conT ''Ord) (return ordType))
[comparepFunction 'compare comparepName (changeTVFlags specifiedSpec bndrs) cons]

ordpFunction comparepName boundVars cons
= comparepFunction comparepName comparepName boundVars cons

comparepFunction compareName comparepName boundVars cons
| null cons = funD compareName [clause [] (normalB [| \x y -> seq x (seq y undefined) |]) []]
| otherwise = funD compareName (concatMap comparepClauses cons)
where
comparepClauses con =
[ mainClause con
, clause [recP conName [], wildP] (normalB [| LT |]) []
, clause [wildP, recP conName []] (normalB [| GT |]) []
] where conName = nameOfCon con

needsOrdP argType con = any ((`occursInType` argType) . nameOfBinder) (boundVars ++ varsBoundInCon con)

mainClause con = do
let conName = nameOfCon con
argTypes = argTypesOfCon con
nArgs = length argTypes

lArgNames <- replicateM nArgs (newName "x")
rArgNames <- replicateM nArgs (newName "y")

clause [ conP conName (map varP lArgNames)
, conP conName (map varP rArgNames)
]
( normalB $ foldr (\(lArg, rArg, argType) rest ->
[| case $(if needsOrdP argType con
then [| $(varE comparepName) $(varE lArg) $(varE rArg) |]
else [| compare $(varE lArg) $(varE rArg) |]) of
EQ -> $rest
o -> o
|])
[| EQ |]
(zip3 lArgNames rArgNames argTypes)
) []

gcompareFunction boundVars cons
| null cons = funD 'gcompare [clause [] (normalB [| \x y -> seq x (seq y undefined) |]) []]
| otherwise = funD 'gcompare (concatMap gcompareClauses cons)
Expand Down
Loading