{-# LANGUAGE AllowAmbiguousTypes #-}
{-# LANGUAGE DataKinds #-}
{-# LANGUAGE FlexibleContexts #-}
{-# LANGUAGE ScopedTypeVariables #-}
{-# LANGUAGE TypeApplications #-}
{-# LANGUAGE TypeFamilies #-}
{-# LANGUAGE TypeOperators #-}

{- | Typed column transformations: the @apply@\/@derive@ family and
default-valued inserts, plus the horizontal merge @('|||')@. Schema changes are
tracked at the type level (e.g. 'applyColumn' rewrites a column's element type
via 'SetColumnType').
-}
module DataFrame.Typed.Apply (
    applyColumn,
    applyMany,
    applyWhere,
    applyAtIndex,
    safeApply,
    deriveWithExpr,
    insertWithDefault,
    insertVectorWithDefault,
    insertUnboxedVector,
    (|||),
) where

import Data.Proxy (Proxy (..))
import qualified Data.Text as T
import qualified Data.Vector as V
import qualified Data.Vector.Unboxed as VU
import GHC.TypeLits (KnownSymbol, Symbol, symbolVal)

import DataFrame.Errors (DataFrameException)
import DataFrame.Internal.Column (Columnable)
import qualified DataFrame.Operations.Core as D
import qualified DataFrame.Operations.Merge as D
import qualified DataFrame.Operations.Transformations as D
import DataFrame.Typed.Freeze (unsafeFreeze)
import DataFrame.Typed.Schema (
    AllKnownSymbol,
    Append,
    AssertAbsent,
    AssertAllColumnsHaveType,
    AssertDisjoint,
    AssertPresent,
    SafeLookup,
    SetColumnType,
    Snoc,
    symbolVals,
 )
import DataFrame.Typed.Types (TExpr (..), TypedDataFrame (..))

{- | Map a function over a column, rewriting its element type from @a@ to @b@.
The schema's entry for @name@ is updated via 'SetColumnType'.

@
df' = applyColumn \@\"age\" (show :: Int -> String) df
-- the \"age\" column is now String-typed
@
-}
applyColumn ::
    forall name a b cols.
    ( KnownSymbol name
    , a ~ SafeLookup name cols
    , Columnable a
    , Columnable b
    , AssertPresent name cols
    ) =>
    (a -> b) ->
    TypedDataFrame cols ->
    TypedDataFrame (SetColumnType name b cols)
applyColumn :: forall (name :: Symbol) a b (cols :: [(Symbol, *)]).
(KnownSymbol name, a ~ SafeLookup name cols, Columnable a,
 Columnable b, AssertPresent name cols) =>
(a -> b)
-> TypedDataFrame cols
-> TypedDataFrame (SetColumnType name b cols)
applyColumn a -> b
f (TDF DataFrame
df) = DataFrame -> TypedDataFrame (SetColumnType name b cols)
forall (cols :: [(Symbol, *)]). DataFrame -> TypedDataFrame cols
unsafeFreeze ((a -> b) -> Text -> DataFrame -> DataFrame
forall b c.
(Columnable b, Columnable c) =>
(b -> c) -> Text -> DataFrame -> DataFrame
D.apply a -> b
f Text
colName DataFrame
df)
  where
    colName :: Text
colName = 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))

-- | Like 'applyColumn' but returns the error instead of throwing.
safeApply ::
    forall name a b cols.
    ( KnownSymbol name
    , a ~ SafeLookup name cols
    , Columnable a
    , Columnable b
    , AssertPresent name cols
    ) =>
    (a -> b) ->
    TypedDataFrame cols ->
    Either DataFrameException (TypedDataFrame (SetColumnType name b cols))
safeApply :: forall (name :: Symbol) a b (cols :: [(Symbol, *)]).
(KnownSymbol name, a ~ SafeLookup name cols, Columnable a,
 Columnable b, AssertPresent name cols) =>
(a -> b)
-> TypedDataFrame cols
-> Either
     DataFrameException (TypedDataFrame (SetColumnType name b cols))
safeApply a -> b
f (TDF DataFrame
df) = (DataFrame -> TypedDataFrame (SetColumnType name b cols))
-> Either DataFrameException DataFrame
-> Either
     DataFrameException (TypedDataFrame (SetColumnType name b cols))
forall a b.
(a -> b)
-> Either DataFrameException a -> Either DataFrameException b
forall (f :: * -> *) a b. Functor f => (a -> b) -> f a -> f b
fmap DataFrame -> TypedDataFrame (SetColumnType name b cols)
forall (cols :: [(Symbol, *)]). DataFrame -> TypedDataFrame cols
unsafeFreeze ((a -> b)
-> Text -> DataFrame -> Either DataFrameException DataFrame
forall b c.
(Columnable b, Columnable c) =>
(b -> c)
-> Text -> DataFrame -> Either DataFrameException DataFrame
D.safeApply a -> b
f Text
colName DataFrame
df)
  where
    colName :: Text
colName = 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))

{- | Apply a type-preserving function to several columns at once. Every named
column must already share the element type @a@ (enforced by
'AssertAllColumnsHaveType').
-}
applyMany ::
    forall (names :: [Symbol]) a cols.
    (AllKnownSymbol names, Columnable a, AssertAllColumnsHaveType names a cols) =>
    (a -> a) ->
    TypedDataFrame cols ->
    TypedDataFrame cols
applyMany :: forall (names :: [Symbol]) a (cols :: [(Symbol, *)]).
(AllKnownSymbol names, Columnable a,
 AssertAllColumnsHaveType names a cols) =>
(a -> a) -> TypedDataFrame cols -> TypedDataFrame cols
applyMany a -> a
f (TDF DataFrame
df) = DataFrame -> TypedDataFrame cols
forall (cols :: [(Symbol, *)]). DataFrame -> TypedDataFrame cols
TDF ((a -> a) -> [Text] -> DataFrame -> DataFrame
forall b c.
(Columnable b, Columnable c) =>
(b -> c) -> [Text] -> DataFrame -> DataFrame
D.applyMany a -> a
f (forall (names :: [Symbol]). AllKnownSymbol names => [Text]
symbolVals @names) DataFrame
df)

{- | Apply a function to a target column only on rows where a condition holds on
a filter column. Both columns are named by type application; the target keeps
its type.

@
applyWhere \@\"flagged\" \@\"score\" id (* 2) df
@
-}
applyWhere ::
    forall filterName targetName a b cols.
    ( KnownSymbol filterName
    , KnownSymbol targetName
    , a ~ SafeLookup filterName cols
    , b ~ SafeLookup targetName cols
    , Columnable a
    , Columnable b
    , AssertPresent filterName cols
    , AssertPresent targetName cols
    ) =>
    (a -> Bool) ->
    (b -> b) ->
    TypedDataFrame cols ->
    TypedDataFrame cols
applyWhere :: forall (filterName :: Symbol) (targetName :: Symbol) a b
       (cols :: [(Symbol, *)]).
(KnownSymbol filterName, KnownSymbol targetName,
 a ~ SafeLookup filterName cols, b ~ SafeLookup targetName cols,
 Columnable a, Columnable b, AssertPresent filterName cols,
 AssertPresent targetName cols) =>
(a -> Bool)
-> (b -> b) -> TypedDataFrame cols -> TypedDataFrame cols
applyWhere a -> Bool
cond b -> b
f (TDF DataFrame
df) = DataFrame -> TypedDataFrame cols
forall (cols :: [(Symbol, *)]). DataFrame -> TypedDataFrame cols
TDF ((a -> Bool) -> Text -> (b -> b) -> Text -> DataFrame -> DataFrame
forall a b.
(Columnable a, Columnable b) =>
(a -> Bool) -> Text -> (b -> b) -> Text -> DataFrame -> DataFrame
D.applyWhere a -> Bool
cond Text
filterName b -> b
f Text
targetName DataFrame
df)
  where
    filterName :: Text
filterName = String -> Text
T.pack (Proxy filterName -> String
forall (n :: Symbol) (proxy :: Symbol -> *).
KnownSymbol n =>
proxy n -> String
symbolVal (forall {k} (t :: k). Proxy t
forall (t :: Symbol). Proxy t
Proxy @filterName))
    targetName :: Text
targetName = String -> Text
T.pack (Proxy targetName -> String
forall (n :: Symbol) (proxy :: Symbol -> *).
KnownSymbol n =>
proxy n -> String
symbolVal (forall {k} (t :: k). Proxy t
forall (t :: Symbol). Proxy t
Proxy @targetName))

-- | Apply a type-preserving function to a single row of a column.
applyAtIndex ::
    forall name a cols.
    ( KnownSymbol name
    , a ~ SafeLookup name cols
    , Columnable a
    , AssertPresent name cols
    ) =>
    Int ->
    (a -> a) ->
    TypedDataFrame cols ->
    TypedDataFrame cols
applyAtIndex :: forall (name :: Symbol) a (cols :: [(Symbol, *)]).
(KnownSymbol name, a ~ SafeLookup name cols, Columnable a,
 AssertPresent name cols) =>
Int -> (a -> a) -> TypedDataFrame cols -> TypedDataFrame cols
applyAtIndex Int
i a -> a
f (TDF DataFrame
df) = DataFrame -> TypedDataFrame cols
forall (cols :: [(Symbol, *)]). DataFrame -> TypedDataFrame cols
TDF (Int -> (a -> a) -> Text -> DataFrame -> DataFrame
forall a.
Columnable a =>
Int -> (a -> a) -> Text -> DataFrame -> DataFrame
D.applyAtIndex Int
i a -> a
f Text
colName DataFrame
df)
  where
    colName :: Text
colName = 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))

{- | Derive a new column and also return a typed reference to it. The returned
expression lives in the extended schema, so it can feed later operations.
-}
deriveWithExpr ::
    forall name a cols.
    ( KnownSymbol name
    , Columnable a
    , AssertAbsent name cols
    ) =>
    TExpr cols a ->
    TypedDataFrame cols ->
    ( TExpr (Snoc cols '(name, a)) a
    , TypedDataFrame (Snoc cols '(name, a))
    )
deriveWithExpr :: forall (name :: Symbol) a (cols :: [(Symbol, *)]).
(KnownSymbol name, Columnable a, AssertAbsent name cols) =>
TExpr cols a
-> TypedDataFrame cols
-> (TExpr (Snoc cols '(name, a)) a,
    TypedDataFrame (Snoc cols '(name, a)))
deriveWithExpr (TExpr Expr a
expr) (TDF DataFrame
df) =
    let (Expr a
e', DataFrame
df') = Text -> Expr a -> DataFrame -> (Expr a, DataFrame)
forall a.
Columnable a =>
Text -> Expr a -> DataFrame -> (Expr a, DataFrame)
D.deriveWithExpr Text
colName Expr a
expr DataFrame
df
     in (Expr a -> TExpr (Snoc cols '(name, a)) a
forall (cols :: [(Symbol, *)]) a. Expr a -> TExpr cols a
TExpr Expr a
e', DataFrame -> TypedDataFrame (Snoc cols '(name, a))
forall (cols :: [(Symbol, *)]). DataFrame -> TypedDataFrame cols
unsafeFreeze DataFrame
df')
  where
    colName :: Text
colName = 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))

-- | Insert a column from a 'Foldable', padding missing rows with a default.
insertWithDefault ::
    forall name a cols t.
    ( KnownSymbol name
    , Columnable a
    , Foldable t
    , AssertAbsent name cols
    ) =>
    a -> t a -> TypedDataFrame cols -> TypedDataFrame ('(name, a) ': cols)
insertWithDefault :: forall (name :: Symbol) a (cols :: [(Symbol, *)]) (t :: * -> *).
(KnownSymbol name, Columnable a, Foldable t,
 AssertAbsent name cols) =>
a
-> t a -> TypedDataFrame cols -> TypedDataFrame ('(name, a) : cols)
insertWithDefault a
def t a
xs (TDF DataFrame
df) =
    DataFrame -> TypedDataFrame ('(name, a) : cols)
forall (cols :: [(Symbol, *)]). DataFrame -> TypedDataFrame cols
unsafeFreeze (a -> Text -> t a -> DataFrame -> DataFrame
forall a (t :: * -> *).
(Columnable a, Foldable t) =>
a -> Text -> t a -> DataFrame -> DataFrame
D.insertWithDefault a
def Text
colName t a
xs DataFrame
df)
  where
    colName :: Text
colName = 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))

-- | Insert a boxed 'V.Vector', padding missing rows with a default.
insertVectorWithDefault ::
    forall name a cols.
    ( KnownSymbol name
    , Columnable a
    , AssertAbsent name cols
    ) =>
    a -> V.Vector a -> TypedDataFrame cols -> TypedDataFrame ('(name, a) ': cols)
insertVectorWithDefault :: forall (name :: Symbol) a (cols :: [(Symbol, *)]).
(KnownSymbol name, Columnable a, AssertAbsent name cols) =>
a
-> Vector a
-> TypedDataFrame cols
-> TypedDataFrame ('(name, a) : cols)
insertVectorWithDefault a
def Vector a
vec (TDF DataFrame
df) =
    DataFrame -> TypedDataFrame ('(name, a) : cols)
forall (cols :: [(Symbol, *)]). DataFrame -> TypedDataFrame cols
unsafeFreeze (a -> Text -> Vector a -> DataFrame -> DataFrame
forall a.
Columnable a =>
a -> Text -> Vector a -> DataFrame -> DataFrame
D.insertVectorWithDefault a
def Text
colName Vector a
vec DataFrame
df)
  where
    colName :: Text
colName = 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))

-- | Insert an unboxed 'VU.Vector' as a new column.
insertUnboxedVector ::
    forall name a cols.
    ( KnownSymbol name
    , Columnable a
    , VU.Unbox a
    , AssertAbsent name cols
    ) =>
    VU.Vector a -> TypedDataFrame cols -> TypedDataFrame ('(name, a) ': cols)
insertUnboxedVector :: forall (name :: Symbol) a (cols :: [(Symbol, *)]).
(KnownSymbol name, Columnable a, Unbox a,
 AssertAbsent name cols) =>
Vector a
-> TypedDataFrame cols -> TypedDataFrame ('(name, a) : cols)
insertUnboxedVector Vector a
vec (TDF DataFrame
df) =
    DataFrame -> TypedDataFrame ('(name, a) : cols)
forall (cols :: [(Symbol, *)]). DataFrame -> TypedDataFrame cols
unsafeFreeze (Text -> Vector a -> DataFrame -> DataFrame
forall a.
(Columnable a, Unbox a) =>
Text -> Vector a -> DataFrame -> DataFrame
D.insertUnboxedVector Text
colName Vector a
vec DataFrame
df)
  where
    colName :: Text
colName = 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))

{- | Horizontal merge: place two DataFrames side by side. The schemas must be
disjoint (no shared column names), enforced by 'AssertDisjoint'; the result
schema is their concatenation.
-}
(|||) ::
    (AssertDisjoint left right) =>
    TypedDataFrame left ->
    TypedDataFrame right ->
    TypedDataFrame (Append left right)
(TDF DataFrame
a) ||| :: forall (left :: [(Symbol, *)]) (right :: [(Symbol, *)]).
AssertDisjoint left right =>
TypedDataFrame left
-> TypedDataFrame right -> TypedDataFrame (Append left right)
||| (TDF DataFrame
b) = DataFrame -> TypedDataFrame (Append left right)
forall (cols :: [(Symbol, *)]). DataFrame -> TypedDataFrame cols
unsafeFreeze (DataFrame
a DataFrame -> DataFrame -> DataFrame
D.||| DataFrame
b)