{-# LANGUAGE AllowAmbiguousTypes #-}
{-# LANGUAGE ConstraintKinds #-}
{-# LANGUAGE DataKinds #-}
{-# LANGUAGE FlexibleContexts #-}
{-# LANGUAGE FlexibleInstances #-}
{-# LANGUAGE GADTs #-}
{-# LANGUAGE MultiParamTypeClasses #-}
{-# LANGUAGE PolyKinds #-}
{-# LANGUAGE RankNTypes #-}
{-# LANGUAGE ScopedTypeVariables #-}
{-# LANGUAGE TypeApplications #-}
{-# LANGUAGE TypeFamilies #-}
{-# LANGUAGE TypeOperators #-}
{-# LANGUAGE UndecidableInstances #-}
module DataFrame.Typed.Schema (
Lookup,
SafeLookup,
HasName,
RemoveColumn,
Impute,
SetColumnType,
SubsetSchema,
ExcludeSchema,
RenameInSchema,
RenameManyInSchema,
Append,
Snoc,
Reverse,
ColumnNames,
AssertAbsent,
AssertPresent,
AssertAllPresent,
AssertKeyTypesMatch,
AssertDisjoint,
AssertAllColumnsHaveType,
AssertRealColumn,
AllColumnsReal,
AllDouble,
IsRealType,
IsElem,
StripAllMaybe,
StripMaybeAt,
SharedNames,
UniqueLeft,
InnerJoinSchema,
LeftJoinSchema,
RightJoinSchema,
FullOuterJoinSchema,
ToMaybe,
WrapMaybe,
WrapMaybeColumns,
CollidingColumns,
GroupKeyColumns,
KnownSchema (..),
schemaColumnNames,
AllKnownSymbol (..),
) where
import Data.Int (Int16, Int32, Int64, Int8)
import Data.Kind (Constraint, Type)
import Data.Proxy (Proxy (..))
import qualified Data.Text as T
import qualified Data.Vector.Unboxed as VU
import Data.Word (Word16, Word32, Word64, Word8)
import GHC.TypeLits
import Type.Reflection (SomeTypeRep, Typeable, someTypeRep)
import DataFrame.Internal.Column (Columnable)
import DataFrame.Internal.Types (These)
type family Lookup (name :: Symbol) (cols :: [(Symbol, Type)]) :: Type where
Lookup name ('(name, a) ': _) = a
Lookup name (_ ': rest) = Lookup name rest
Lookup name '[] =
TypeError
('Text "Column '" ':<>: 'Text name ':<>: 'Text "' not found in schema")
type family SafeLookup (name :: Symbol) (cols :: [(Symbol, Type)]) :: Type where
SafeLookup name ('(name, a) ': _) = a
SafeLookup name (_ ': rest) = SafeLookup name rest
SafeLookup name '[] = Int
type family Impute (name :: Symbol) (cols :: [(Symbol, Type)]) :: [(Symbol, Type)] where
Impute name ('(name, Maybe a) ': rest) = '(name, a) ': rest
Impute name ('(name, _) ': rest) =
TypeError
('Text "Column '" ':<>: 'Text name ':<>: 'Text "' is not of kind Maybe *")
Impute name (col ': rest) = col ': Impute name rest
Impute name '[] = '[]
type family
SetColumnType (name :: Symbol) (b :: Type) (cols :: [(Symbol, Type)]) ::
[(Symbol, Type)]
where
SetColumnType name b ('(name, _) ': rest) = '(name, b) ': rest
SetColumnType name b (col ': rest) = col ': SetColumnType name b rest
SetColumnType name b '[] =
TypeError
('Text "Column '" ':<>: 'Text name ':<>: 'Text "' not found in schema")
type family Snoc (xs :: [k]) (x :: k) :: [k] where
Snoc '[] x = '[x]
Snoc (y ': ys) x = y ': Snoc ys x
type family HasName (name :: Symbol) (cols :: [(Symbol, Type)]) :: Bool where
HasName name ('(name, _) ': _) = 'True
HasName name (_ ': rest) = HasName name rest
HasName name '[] = 'False
type family RemoveColumn (name :: Symbol) (cols :: [(Symbol, Type)]) :: [(Symbol, Type)] where
RemoveColumn name ('(name, _) ': rest) = rest
RemoveColumn name (col ': rest) = col ': RemoveColumn name rest
RemoveColumn name '[] = '[]
type family SubsetSchema (names :: [Symbol]) (cols :: [(Symbol, Type)]) :: [(Symbol, Type)] where
SubsetSchema '[] cols = '[]
SubsetSchema (n ': ns) cols = '(n, Lookup n cols) ': SubsetSchema ns cols
type family ExcludeSchema (names :: [Symbol]) (cols :: [(Symbol, Type)]) :: [(Symbol, Type)] where
ExcludeSchema names '[] = '[]
ExcludeSchema names ('(n, a) ': rest) =
ExcludeSchemaHelper (IsElem n names) n a names rest
type family
ExcludeSchemaHelper
(found :: Bool)
(n :: Symbol)
(a :: Type)
(names :: [Symbol])
(rest :: [(Symbol, Type)]) ::
[(Symbol, Type)]
where
ExcludeSchemaHelper 'True n a names rest = ExcludeSchema names rest
ExcludeSchemaHelper 'False n a names rest =
'(n, a) ': ExcludeSchema names rest
type family IsElem (x :: Symbol) (xs :: [Symbol]) :: Bool where
IsElem x '[] = 'False
IsElem x (x ': _) = 'True
IsElem x (_ ': xs) = IsElem x xs
type family
RenameInSchema (old :: Symbol) (new :: Symbol) (cols :: [(Symbol, Type)]) ::
[(Symbol, Type)]
where
RenameInSchema old new ('(old, a) ': rest) = '(new, a) ': rest
RenameInSchema old new (col ': rest) = col ': RenameInSchema old new rest
RenameInSchema old new '[] =
TypeError
('Text "Cannot rename: column '" ':<>: 'Text old ':<>: 'Text "' not found")
type family
RenameManyInSchema (pairs :: [(Symbol, Symbol)]) (cols :: [(Symbol, Type)]) ::
[(Symbol, Type)]
where
RenameManyInSchema '[] cols = cols
RenameManyInSchema ('(old, new) ': rest) cols =
RenameManyInSchema rest (RenameInSchema old new cols)
type family Append (xs :: [k]) (ys :: [k]) :: [k] where
Append '[] ys = ys
Append (x ': xs) ys = x ': Append xs ys
type family Reverse (xs :: [(Symbol, Type)]) :: [(Symbol, Type)] where
Reverse xs = ReverseAcc xs '[]
type family
ReverseAcc (xs :: [(Symbol, Type)]) (acc :: [(Symbol, Type)]) ::
[(Symbol, Type)]
where
ReverseAcc '[] acc = acc
ReverseAcc (x ': xs) acc = ReverseAcc xs (x ': acc)
type family ColumnNames (cols :: [(Symbol, Type)]) :: [Symbol] where
ColumnNames '[] = '[]
ColumnNames ('(n, _) ': rest) = n ': ColumnNames rest
type family AssertAbsent (name :: Symbol) (cols :: [(Symbol, Type)]) :: Constraint where
AssertAbsent name cols = AssertAbsentHelper name (HasName name cols) cols
type family
AssertAbsentHelper (name :: Symbol) (found :: Bool) (cols :: [(Symbol, Type)]) ::
Constraint
where
AssertAbsentHelper name 'False cols = ()
AssertAbsentHelper name 'True cols =
TypeError
( 'Text "Column '"
':<>: 'Text name
':<>: 'Text "' already exists in schema. "
':<>: 'Text "Use replaceColumn to overwrite."
)
type family AssertPresent (name :: Symbol) (cols :: [(Symbol, Type)]) :: Constraint where
AssertPresent name cols = AssertPresentHelper name (HasName name cols) cols
type family
AssertPresentHelper (name :: Symbol) (found :: Bool) (cols :: [(Symbol, Type)]) ::
Constraint
where
AssertPresentHelper name 'True cols = ()
AssertPresentHelper name 'False cols =
TypeError
('Text "Column '" ':<>: 'Text name ':<>: 'Text "' not found in schema")
type family AssertAllPresent (name :: [Symbol]) (cols :: [(Symbol, Type)]) :: Constraint where
AssertAllPresent (name ': rest) cols =
AssertAllPresentHelper (HasName name cols) name rest cols
AssertAllPresent '[] cols = ()
type family
AssertAllPresentHelper
(found :: Bool)
(name :: Symbol)
(rest :: [Symbol])
(cols :: [(Symbol, Type)]) ::
Constraint
where
AssertAllPresentHelper 'True name rest cols = AssertAllPresent rest cols
AssertAllPresentHelper 'False name rest cols =
TypeError
('Text "Column '" ':<>: 'Text name ':<>: 'Text "' not found in schema")
type family
AssertKeyTypesMatch
(keys :: [Symbol])
(left :: [(Symbol, Type)])
(right :: [(Symbol, Type)]) ::
Constraint
where
AssertKeyTypesMatch '[] left right = ()
AssertKeyTypesMatch (k ': ks) left right =
( KeyTypeMatchHelper k (SafeLookup k left) (SafeLookup k right)
, AssertKeyTypesMatch ks left right
)
type family
KeyTypeMatchHelper (k :: Symbol) (l :: Type) (r :: Type) ::
Constraint
where
KeyTypeMatchHelper k a a = ()
KeyTypeMatchHelper k (Maybe a) a = ()
KeyTypeMatchHelper k a (Maybe a) = ()
KeyTypeMatchHelper k l r =
TypeError
( 'Text "Join key '"
':<>: 'Text k
':<>: 'Text "' has type "
':<>: 'ShowType l
':<>: 'Text " in the left table but "
':<>: 'ShowType r
':<>: 'Text " in the right table"
)
type family
AssertDisjoint (left :: [(Symbol, Type)]) (right :: [(Symbol, Type)]) ::
Constraint
where
AssertDisjoint left right =
AssertDisjointHelper (SharedNames left right) left right
type family
AssertDisjointHelper
(shared :: [Symbol])
(left :: [(Symbol, Type)])
(right :: [(Symbol, Type)]) ::
Constraint
where
AssertDisjointHelper '[] left right = ()
AssertDisjointHelper (n ': ns) left right =
TypeError
( 'Text "Cannot horizontally merge: column '"
':<>: 'Text n
':<>: 'Text "' appears in both schemas"
)
type family
AssertAllColumnsHaveType
(names :: [Symbol])
(a :: Type)
(cols :: [(Symbol, Type)]) ::
Constraint
where
AssertAllColumnsHaveType '[] a cols = ()
AssertAllColumnsHaveType (n ': ns) a cols =
( SafeLookup n cols ~ a
, AssertPresent n cols
, AssertAllColumnsHaveType ns a cols
)
type family IsRealType (a :: Type) :: Bool where
IsRealType Int = 'True
IsRealType Int8 = 'True
IsRealType Int16 = 'True
IsRealType Int32 = 'True
IsRealType Int64 = 'True
IsRealType Word = 'True
IsRealType Word8 = 'True
IsRealType Word16 = 'True
IsRealType Word32 = 'True
IsRealType Word64 = 'True
IsRealType Double = 'True
IsRealType Float = 'True
IsRealType _ = 'False
type family AssertRealColumn (fn :: Symbol) (name :: Symbol) (a :: Type) :: Constraint where
AssertRealColumn fn name a = AssertRealColumnGo fn name a (IsRealType a)
type family
AssertRealColumnGo (fn :: Symbol) (name :: Symbol) (a :: Type) (isReal :: Bool) ::
Constraint
where
AssertRealColumnGo fn name a 'True = ()
AssertRealColumnGo fn name a 'False =
TypeError
( 'Text fn
':<>: 'Text ": expected a real number column for '"
':<>: 'Text name
':<>: 'Text "' but instead you gave "
':<>: 'ShowType a
)
type family AllColumnsReal (fn :: Symbol) (cols :: [(Symbol, Type)]) :: Constraint where
AllColumnsReal fn '[] = ()
AllColumnsReal fn ('(n, a) ': rest) =
(AssertRealColumn fn n a, Real a, VU.Unbox a, AllColumnsReal fn rest)
type family AllDouble (cols :: [(Symbol, Type)]) :: Constraint where
AllDouble '[] = ()
AllDouble ('(n, Double) ': rest) = AllDouble rest
AllDouble ('(n, a) ': rest) =
TypeError
( 'Text "Column '"
':<>: 'Text n
':<>: 'Text "' must be Double for this model, but is "
':<>: 'ShowType a
':$$: 'Text "Convert it (toDouble) or drop it before fitting."
)
type family StripAllMaybe (cols :: [(Symbol, Type)]) :: [(Symbol, Type)] where
StripAllMaybe '[] = '[]
StripAllMaybe ('(n, Maybe a) ': rest) = '(n, a) ': StripAllMaybe rest
StripAllMaybe ('(n, a) ': rest) = '(n, a) ': StripAllMaybe rest
type family StripMaybeAt (name :: Symbol) (cols :: [(Symbol, Type)]) :: [(Symbol, Type)] where
StripMaybeAt name ('(name, Maybe a) ': rest) = '(name, a) ': rest
StripMaybeAt name ('(name, a) ': rest) = '(name, a) ': rest
StripMaybeAt name (col ': rest) = col ': StripMaybeAt name rest
StripMaybeAt name '[] =
TypeError
('Text "Column '" ':<>: 'Text name ':<>: 'Text "' not found in schema")
type family SharedNames (left :: [(Symbol, Type)]) (right :: [(Symbol, Type)]) :: [Symbol] where
SharedNames '[] right = '[]
SharedNames ('(n, _) ': rest) right =
SharedNamesHelper (HasName n right) n rest right
type family
SharedNamesHelper
(found :: Bool)
(n :: Symbol)
(rest :: [(Symbol, Type)])
(right :: [(Symbol, Type)]) ::
[Symbol]
where
SharedNamesHelper 'True n rest right = n ': SharedNames rest right
SharedNamesHelper 'False n rest right = SharedNames rest right
type family
UniqueLeft (left :: [(Symbol, Type)]) (rightNames :: [Symbol]) ::
[(Symbol, Type)]
where
UniqueLeft '[] _ = '[]
UniqueLeft ('(n, a) ': rest) rn =
UniqueLeftHelper (IsElem n rn) n a rest rn
type family
UniqueLeftHelper
(found :: Bool)
(n :: Symbol)
(a :: Type)
(rest :: [(Symbol, Type)])
(rn :: [Symbol]) ::
[(Symbol, Type)]
where
UniqueLeftHelper 'True n a rest rn = UniqueLeft rest rn
UniqueLeftHelper 'False n a rest rn = '(n, a) ': UniqueLeft rest rn
type family ToMaybe (a :: Type) :: Type where
ToMaybe (Maybe a) = Maybe a
ToMaybe a = Maybe a
type family WrapMaybe (cols :: [(Symbol, Type)]) :: [(Symbol, Type)] where
WrapMaybe '[] = '[]
WrapMaybe ('(n, a) ': rest) = '(n, ToMaybe a) ': WrapMaybe rest
type family
WrapMaybeColumns (names :: [Symbol]) (cols :: [(Symbol, Type)]) ::
[(Symbol, Type)]
where
WrapMaybeColumns names '[] = '[]
WrapMaybeColumns names ('(n, a) ': rest) =
WrapMaybeColumnsHelper (IsElem n names) n a names rest
type family
WrapMaybeColumnsHelper
(found :: Bool)
(n :: Symbol)
(a :: Type)
(names :: [Symbol])
(rest :: [(Symbol, Type)]) ::
[(Symbol, Type)]
where
WrapMaybeColumnsHelper 'True n a names rest =
'(n, ToMaybe a) ': WrapMaybeColumns names rest
WrapMaybeColumnsHelper 'False n a names rest =
'(n, a) ': WrapMaybeColumns names rest
type family
CollidingColumns
(left :: [(Symbol, Type)])
(right :: [(Symbol, Type)])
(keys :: [Symbol]) ::
[(Symbol, Type)]
where
CollidingColumns '[] _ _ = '[]
CollidingColumns ('(n, a) ': rest) right keys =
CollidingColumnsHelper1 (IsElem n keys) n a rest right keys
type family
CollidingColumnsHelper1
(isKey :: Bool)
(n :: Symbol)
(a :: Type)
(rest :: [(Symbol, Type)])
(right :: [(Symbol, Type)])
(keys :: [Symbol]) ::
[(Symbol, Type)]
where
CollidingColumnsHelper1 'True n a rest right keys =
CollidingColumns rest right keys
CollidingColumnsHelper1 'False n a rest right keys =
CollidingColumnsHelper2 (HasName n right) n a rest right keys
type family
CollidingColumnsHelper2
(inRight :: Bool)
(n :: Symbol)
(a :: Type)
(rest :: [(Symbol, Type)])
(right :: [(Symbol, Type)])
(keys :: [Symbol]) ::
[(Symbol, Type)]
where
CollidingColumnsHelper2 'True n a rest right keys =
'(n, These a (Lookup n right)) ': CollidingColumns rest right keys
CollidingColumnsHelper2 'False n a rest right keys =
CollidingColumns rest right keys
type family
InnerJoinSchema
(keys :: [Symbol])
(left :: [(Symbol, Type)])
(right :: [(Symbol, Type)]) ::
[(Symbol, Type)]
where
InnerJoinSchema keys left right =
Append
(SubsetSchema keys left)
( Append
(UniqueLeft left (Append keys (ColumnNames right)))
( Append
(UniqueLeft right (Append keys (ColumnNames left)))
(CollidingColumns left right keys)
)
)
type family
LeftJoinSchema
(keys :: [Symbol])
(left :: [(Symbol, Type)])
(right :: [(Symbol, Type)]) ::
[(Symbol, Type)]
where
LeftJoinSchema keys left right =
Append
(SubsetSchema keys left)
( Append
(UniqueLeft left (Append keys (ColumnNames right)))
( Append
(WrapMaybe (UniqueLeft right (Append keys (ColumnNames left))))
(CollidingColumns left right keys)
)
)
type family
RightJoinSchema
(keys :: [Symbol])
(left :: [(Symbol, Type)])
(right :: [(Symbol, Type)]) ::
[(Symbol, Type)]
where
RightJoinSchema keys left right =
Append
(SubsetSchema keys right)
( Append
(WrapMaybe (UniqueLeft left (Append keys (ColumnNames right))))
( Append
(UniqueLeft right (Append keys (ColumnNames left)))
(CollidingColumns left right keys)
)
)
type family
FullOuterJoinSchema
(keys :: [Symbol])
(left :: [(Symbol, Type)])
(right :: [(Symbol, Type)]) ::
[(Symbol, Type)]
where
FullOuterJoinSchema keys left right =
Append
(WrapMaybe (SubsetSchema keys left))
( Append
(WrapMaybe (UniqueLeft left (Append keys (ColumnNames right))))
( Append
(WrapMaybe (UniqueLeft right (Append keys (ColumnNames left))))
(CollidingColumns left right keys)
)
)
type family
GroupKeyColumns (keys :: [Symbol]) (cols :: [(Symbol, Type)]) ::
[(Symbol, Type)]
where
GroupKeyColumns keys '[] = '[]
GroupKeyColumns keys ('(n, a) ': rest) =
GroupKeyColumnsHelper (IsElem n keys) n a keys rest
type family
GroupKeyColumnsHelper
(found :: Bool)
(n :: Symbol)
(a :: Type)
(keys :: [Symbol])
(rest :: [(Symbol, Type)]) ::
[(Symbol, Type)]
where
GroupKeyColumnsHelper 'True n a keys rest =
'(n, a) ': GroupKeyColumns keys rest
GroupKeyColumnsHelper 'False n a keys rest = GroupKeyColumns keys rest
class KnownSchema (cols :: [(Symbol, Type)]) where
schemaEvidence :: [(T.Text, SomeTypeRep)]
instance KnownSchema '[] where
schemaEvidence :: [(Text, SomeTypeRep)]
schemaEvidence = []
instance
(KnownSymbol name, Typeable a, Columnable a, KnownSchema rest) =>
KnownSchema ('(name, a) ': rest)
where
schemaEvidence :: [(Text, SomeTypeRep)]
schemaEvidence =
(String -> Text
T.pack (Proxy name -> String
forall (n :: Symbol) (proxy :: Symbol -> *).
KnownSymbol n =>
proxy n -> String
symbolVal (forall {k} (t :: k). Proxy t
forall (t :: Symbol). Proxy t
Proxy @name)), Proxy a -> SomeTypeRep
forall {k} (proxy :: k -> *) (a :: k).
Typeable a =>
proxy a -> SomeTypeRep
someTypeRep (forall t. Proxy t
forall {k} (t :: k). Proxy t
Proxy @a))
(Text, SomeTypeRep)
-> [(Text, SomeTypeRep)] -> [(Text, SomeTypeRep)]
forall a. a -> [a] -> [a]
: forall (cols :: [(Symbol, *)]).
KnownSchema cols =>
[(Text, SomeTypeRep)]
schemaEvidence @rest
schemaColumnNames :: forall cols. (KnownSchema cols) => [T.Text]
schemaColumnNames :: forall (cols :: [(Symbol, *)]). KnownSchema cols => [Text]
schemaColumnNames = ((Text, SomeTypeRep) -> Text) -> [(Text, SomeTypeRep)] -> [Text]
forall a b. (a -> b) -> [a] -> [b]
map (Text, SomeTypeRep) -> Text
forall a b. (a, b) -> a
fst (forall (cols :: [(Symbol, *)]).
KnownSchema cols =>
[(Text, SomeTypeRep)]
schemaEvidence @cols)
class AllKnownSymbol (names :: [Symbol]) where
symbolVals :: [T.Text]
instance AllKnownSymbol '[] where
symbolVals :: [Text]
symbolVals = []
instance (KnownSymbol n, AllKnownSymbol ns) => AllKnownSymbol (n ': ns) where
symbolVals :: [Text]
symbolVals = String -> Text
T.pack (Proxy n -> String
forall (n :: Symbol) (proxy :: Symbol -> *).
KnownSymbol n =>
proxy n -> String
symbolVal (forall {k} (t :: k). Proxy t
forall (t :: Symbol). Proxy t
Proxy @n)) Text -> [Text] -> [Text]
forall a. a -> [a] -> [a]
: forall (names :: [Symbol]). AllKnownSymbol names => [Text]
symbolVals @ns